mirror of
https://github.com/phishingclub/phishingclub.git
synced 2026-10-08 00:16:54 +02:00
250 files changed
+47469
No files matched your search
+128
@@ -0,0 +1,128 @@
|
||||
package g
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"compress/flate"
|
||||
"compress/gzip"
|
||||
"compress/zlib"
|
||||
"io"
|
||||
)
|
||||
|
||||
type (
|
||||
// A struct that wraps Bytes for compression.
|
||||
bcompress struct{ bytes Bytes }
|
||||
|
||||
// A struct that wraps Bytes for decompression.
|
||||
bdecompress struct{ bytes Bytes }
|
||||
)
|
||||
|
||||
// Compress returns a bcompress struct wrapping the given Bytes.
|
||||
func (bs Bytes) Compress() bcompress { return bcompress{bs} }
|
||||
|
||||
// Decompress returns a bdecompress struct wrapping the given Bytes.
|
||||
func (bs Bytes) Decompress() bdecompress { return bdecompress{bs} }
|
||||
|
||||
// Zlib compresses the wrapped Bytes using the zlib compression algorithm and
|
||||
// returns the compressed data as Bytes.
|
||||
func (c bcompress) Zlib() Bytes {
|
||||
buffer := new(bytes.Buffer)
|
||||
writer := zlib.NewWriter(buffer)
|
||||
|
||||
_, _ = writer.Write(c.bytes)
|
||||
_ = writer.Flush()
|
||||
_ = writer.Close()
|
||||
|
||||
return buffer.Bytes()
|
||||
}
|
||||
|
||||
// Zlib decompresses the wrapped Bytes using the zlib compression algorithm and
|
||||
// returns the decompressed data as a Result[Bytes].
|
||||
func (d bdecompress) Zlib() Result[Bytes] {
|
||||
reader, err := zlib.NewReader(bytes.NewReader(d.bytes))
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
defer reader.Close()
|
||||
|
||||
buffer := new(bytes.Buffer)
|
||||
if _, err := io.Copy(buffer, reader); err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(Bytes(buffer.Bytes()))
|
||||
}
|
||||
|
||||
// Gzip compresses the wrapped Bytes using the gzip compression format and
|
||||
// returns the compressed data as Bytes.
|
||||
func (c bcompress) Gzip() Bytes {
|
||||
buffer := new(bytes.Buffer)
|
||||
writer := gzip.NewWriter(buffer)
|
||||
|
||||
_, _ = writer.Write(c.bytes)
|
||||
_ = writer.Flush()
|
||||
_ = writer.Close()
|
||||
|
||||
return buffer.Bytes()
|
||||
}
|
||||
|
||||
// Gzip decompresses the wrapped Bytes using the gzip compression format and
|
||||
// returns the decompressed data as a Result[Bytes].
|
||||
func (d bdecompress) Gzip() Result[Bytes] {
|
||||
reader, err := gzip.NewReader(bytes.NewReader(d.bytes))
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
defer reader.Close()
|
||||
|
||||
buffer := new(bytes.Buffer)
|
||||
if _, err := io.Copy(buffer, reader); err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(Bytes(buffer.Bytes()))
|
||||
}
|
||||
|
||||
// Flate compresses the wrapped Bytes using the flate (deflate) compression
|
||||
// algorithm and returns the compressed data as Bytes.
|
||||
// It accepts an optional compression level. If no level is provided, it
|
||||
// defaults to 7. The level is clamped to the valid flate range [-2, 9] to
|
||||
// avoid an invalid-level error that would otherwise return a nil writer and
|
||||
// panic on use.
|
||||
func (c bcompress) Flate(level ...int) Bytes {
|
||||
buffer := new(bytes.Buffer)
|
||||
|
||||
l := 7
|
||||
if len(level) != 0 {
|
||||
l = level[0]
|
||||
}
|
||||
|
||||
if l < flate.HuffmanOnly {
|
||||
l = flate.HuffmanOnly
|
||||
} else if l > flate.BestCompression {
|
||||
l = flate.BestCompression
|
||||
}
|
||||
|
||||
writer, _ := flate.NewWriter(buffer, l)
|
||||
|
||||
_, _ = writer.Write(c.bytes)
|
||||
_ = writer.Flush()
|
||||
_ = writer.Close()
|
||||
|
||||
return buffer.Bytes()
|
||||
}
|
||||
|
||||
// Flate decompresses the wrapped Bytes using the flate (deflate) compression
|
||||
// algorithm and returns the decompressed data as a Result[Bytes].
|
||||
func (d bdecompress) Flate() Result[Bytes] {
|
||||
reader := flate.NewReader(bytes.NewReader(d.bytes))
|
||||
defer reader.Close()
|
||||
|
||||
buffer := new(bytes.Buffer)
|
||||
if _, err := io.Copy(buffer, reader); err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(Bytes(buffer.Bytes()))
|
||||
}
|
||||
+344
@@ -0,0 +1,344 @@
|
||||
package g
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"html"
|
||||
"net/url"
|
||||
"strconv"
|
||||
|
||||
json "encoding/json/v2"
|
||||
)
|
||||
|
||||
type (
|
||||
// A struct that wraps Bytes for encoding.
|
||||
bencode struct{ bytes Bytes }
|
||||
|
||||
// A struct that wraps Bytes for decoding.
|
||||
bdecode struct{ bytes Bytes }
|
||||
)
|
||||
|
||||
// Encode returns a bencode struct wrapping the given Bytes.
|
||||
func (bs Bytes) Encode() bencode { return bencode{bs} }
|
||||
|
||||
// Decode returns a bdecode struct wrapping the given Bytes.
|
||||
func (bs Bytes) Decode() bdecode { return bdecode{bs} }
|
||||
|
||||
// Base64 encodes the wrapped Bytes using standard Base64 (with padding).
|
||||
func (e bencode) Base64() Bytes {
|
||||
out := make(Bytes, base64.StdEncoding.EncodedLen(len(e.bytes)))
|
||||
base64.StdEncoding.Encode(out, e.bytes)
|
||||
return out
|
||||
}
|
||||
|
||||
// Base64Raw encodes the wrapped Bytes using standard Base64 without padding.
|
||||
func (e bencode) Base64Raw() Bytes {
|
||||
out := make(Bytes, base64.RawStdEncoding.EncodedLen(len(e.bytes)))
|
||||
base64.RawStdEncoding.Encode(out, e.bytes)
|
||||
return out
|
||||
}
|
||||
|
||||
// Base64URL encodes the wrapped Bytes using URL-safe Base64 (with padding).
|
||||
func (e bencode) Base64URL() Bytes {
|
||||
out := make(Bytes, base64.URLEncoding.EncodedLen(len(e.bytes)))
|
||||
base64.URLEncoding.Encode(out, e.bytes)
|
||||
return out
|
||||
}
|
||||
|
||||
// Base64RawURL encodes the wrapped Bytes using URL-safe Base64 without padding.
|
||||
func (e bencode) Base64RawURL() Bytes {
|
||||
out := make(Bytes, base64.RawURLEncoding.EncodedLen(len(e.bytes)))
|
||||
base64.RawURLEncoding.Encode(out, e.bytes)
|
||||
return out
|
||||
}
|
||||
|
||||
// Base64 decodes the wrapped Bytes as standard Base64 (with padding) and returns Result[Bytes].
|
||||
func (d bdecode) Base64() Result[Bytes] {
|
||||
out := make(Bytes, base64.StdEncoding.DecodedLen(len(d.bytes)))
|
||||
n, err := base64.StdEncoding.Decode(out, d.bytes)
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(out[:n])
|
||||
}
|
||||
|
||||
// Base64Raw decodes the wrapped Bytes as standard Base64 without padding and returns Result[Bytes].
|
||||
func (d bdecode) Base64Raw() Result[Bytes] {
|
||||
out := make(Bytes, base64.RawStdEncoding.DecodedLen(len(d.bytes)))
|
||||
n, err := base64.RawStdEncoding.Decode(out, d.bytes)
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(out[:n])
|
||||
}
|
||||
|
||||
// Base64URL decodes the wrapped Bytes as URL-safe Base64 (with padding) and returns Result[Bytes].
|
||||
func (d bdecode) Base64URL() Result[Bytes] {
|
||||
out := make(Bytes, base64.URLEncoding.DecodedLen(len(d.bytes)))
|
||||
n, err := base64.URLEncoding.Decode(out, d.bytes)
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(out[:n])
|
||||
}
|
||||
|
||||
// Base64RawURL decodes the wrapped Bytes as URL-safe Base64 without padding and returns Result[Bytes].
|
||||
func (d bdecode) Base64RawURL() Result[Bytes] {
|
||||
out := make(Bytes, base64.RawURLEncoding.DecodedLen(len(d.bytes)))
|
||||
n, err := base64.RawURLEncoding.Decode(out, d.bytes)
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(out[:n])
|
||||
}
|
||||
|
||||
// Hex hex-encodes the wrapped Bytes and returns the result as Bytes.
|
||||
func (e bencode) Hex() Bytes {
|
||||
out := make(Bytes, hex.EncodedLen(len(e.bytes)))
|
||||
hex.Encode(out, e.bytes)
|
||||
return out
|
||||
}
|
||||
|
||||
// Hex hex-decodes the wrapped Bytes and returns the decoded result as Result[Bytes].
|
||||
func (d bdecode) Hex() Result[Bytes] {
|
||||
out := make(Bytes, hex.DecodedLen(len(d.bytes)))
|
||||
n, err := hex.Decode(out, d.bytes)
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(out[:n])
|
||||
}
|
||||
|
||||
// XOR encodes the wrapped Bytes using XOR cipher with the given key.
|
||||
//
|
||||
// Warning: a repeating-key XOR cipher is not a security primitive and provides
|
||||
// no real confidentiality. Use it only for lightweight obfuscation, never to
|
||||
// protect secrets.
|
||||
func (e bencode) XOR(key Bytes) Bytes {
|
||||
if len(key) == 0 {
|
||||
return e.bytes.Clone()
|
||||
}
|
||||
|
||||
out := make(Bytes, len(e.bytes))
|
||||
for i, b := range e.bytes {
|
||||
out[i] = b ^ key[i%len(key)]
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
// XOR decodes the wrapped Bytes using XOR cipher with the given key.
|
||||
func (d bdecode) XOR(key Bytes) Bytes { return d.bytes.Encode().XOR(key) }
|
||||
|
||||
// Binary converts the wrapped Bytes to its binary representation as Bytes.
|
||||
func (e bencode) Binary() Bytes {
|
||||
var b Builder
|
||||
b.Grow(e.bytes.Len() * 8)
|
||||
|
||||
for _, c := range e.bytes {
|
||||
for bit := 7; bit >= 0; bit-- {
|
||||
b.WriteByte('0' + (c>>uint(bit))&1)
|
||||
}
|
||||
}
|
||||
|
||||
return b.String().Bytes()
|
||||
}
|
||||
|
||||
// Binary converts the wrapped binary Bytes back to raw Bytes as Result[Bytes].
|
||||
func (d bdecode) Binary() Result[Bytes] {
|
||||
if len(d.bytes)%8 != 0 {
|
||||
return Err[Bytes](ErrInvalidBinaryLength)
|
||||
}
|
||||
|
||||
out := make(Bytes, 0, len(d.bytes)/8)
|
||||
|
||||
for i := 0; i+8 <= len(d.bytes); i += 8 {
|
||||
var b byte
|
||||
for j := range 8 {
|
||||
c := d.bytes[i+j]
|
||||
if c != '0' && c != '1' {
|
||||
return Err[Bytes](ErrInvalidBinaryDigit)
|
||||
}
|
||||
b = b<<1 | (c - '0')
|
||||
}
|
||||
|
||||
out = append(out, b)
|
||||
}
|
||||
|
||||
return Ok(out)
|
||||
}
|
||||
|
||||
// JSON encodes the wrapped Bytes as a JSON string using encoding/json/v2 and
|
||||
// returns the result as Result[Bytes].
|
||||
//
|
||||
// The bytes are treated as text, mirroring String.Encode().JSON.
|
||||
//
|
||||
// v2 semantics: Bytes containing invalid UTF-8 yield Err — use Base64/Hex
|
||||
// encoding for arbitrary binary data.
|
||||
// Unlike encoding/json v1, the output does not HTML-escape '<', '>', '&' or the
|
||||
// line separators U+2028/U+2029 — they are emitted raw. Escape the output
|
||||
// yourself before embedding it in HTML or <script> contexts.
|
||||
func (e bencode) JSON() Result[Bytes] {
|
||||
jsonData, err := json.Marshal(string(e.bytes))
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(Bytes(jsonData))
|
||||
}
|
||||
|
||||
// JSON decodes the wrapped JSON string using encoding/json/v2 and returns the
|
||||
// result as Result[Bytes].
|
||||
//
|
||||
// v2 semantics: a JSON string containing invalid UTF-8 yields Err instead of
|
||||
// being decoded with U+FFFD replacements.
|
||||
func (d bdecode) JSON() Result[Bytes] {
|
||||
var data String
|
||||
err := json.Unmarshal(d.bytes, &data)
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(data.Bytes())
|
||||
}
|
||||
|
||||
// URL encodes the wrapped Bytes, leaving the RFC 2396 reserved characters
|
||||
// (";/?:@&=+$,") unescaped and query-escaping the rest. If safe characters are
|
||||
// provided, they replace that default set and will not be encoded.
|
||||
//
|
||||
// Unlike String.Encode().URL, which matches safe characters by rune, matching
|
||||
// here is byte-wise; with ASCII-only safe sets the output is identical to the
|
||||
// String version for valid UTF-8 input. Non-UTF-8 bytes are percent-encoded
|
||||
// verbatim, so the encoding is lossless for arbitrary binary input.
|
||||
func (e bencode) URL(safe ...Bytes) Bytes {
|
||||
reserved := Bytes(";/?:@&=+$,")
|
||||
if len(safe) != 0 {
|
||||
reserved = safe[0]
|
||||
}
|
||||
|
||||
out := make(Bytes, 0, len(e.bytes))
|
||||
|
||||
for _, c := range e.bytes {
|
||||
if reserved.IndexByte(c) != -1 {
|
||||
out = append(out, c)
|
||||
continue
|
||||
}
|
||||
|
||||
out = appendQueryEscaped(out, c)
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
const upperhex = "0123456789ABCDEF"
|
||||
|
||||
// appendQueryEscaped appends c to dst using url.QueryEscape semantics:
|
||||
// unreserved bytes (A-Z, a-z, 0-9, '-', '_', '.', '~') pass through unchanged,
|
||||
// a space becomes '+', and every other byte is percent-encoded.
|
||||
func appendQueryEscaped(dst Bytes, c byte) Bytes {
|
||||
switch {
|
||||
case 'A' <= c && c <= 'Z', 'a' <= c && c <= 'z', '0' <= c && c <= '9',
|
||||
c == '-', c == '_', c == '.', c == '~':
|
||||
return append(dst, c)
|
||||
case c == ' ':
|
||||
return append(dst, '+')
|
||||
default:
|
||||
return append(dst, '%', upperhex[c>>4], upperhex[c&0xF])
|
||||
}
|
||||
}
|
||||
|
||||
// URL URL-decodes the wrapped Bytes and returns the decoded result as Result[Bytes].
|
||||
func (d bdecode) URL() Result[Bytes] {
|
||||
result, err := url.QueryUnescape(string(d.bytes))
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
return Ok(Bytes(result))
|
||||
}
|
||||
|
||||
// HTML HTML-encodes the wrapped Bytes, escaping the characters <, >, &, ' and ".
|
||||
// All other bytes, including non-UTF-8 sequences, pass through unchanged.
|
||||
func (e bencode) HTML() Bytes { return Bytes(html.EscapeString(string(e.bytes))) }
|
||||
|
||||
// HTML HTML-decodes the wrapped Bytes, unescaping HTML entities.
|
||||
// Bytes that are not part of an entity, including non-UTF-8 sequences, pass through unchanged.
|
||||
func (d bdecode) HTML() Bytes { return Bytes(html.UnescapeString(string(d.bytes))) }
|
||||
|
||||
// Rot13 encodes the wrapped Bytes using the ROT13 cipher.
|
||||
//
|
||||
// The rotation is byte-wise over the ASCII letters A-Z and a-z; all other
|
||||
// bytes, including those of multibyte UTF-8 runes, are left untouched, so the
|
||||
// output is identical to String.Encode().Rot13 for valid UTF-8 input.
|
||||
//
|
||||
// WARNING: ROT13 is NOT a security primitive. It is a fixed letter-substitution
|
||||
// cipher with no key and is trivially reversible. Use it only for obfuscation.
|
||||
func (e bencode) Rot13() Bytes {
|
||||
out := make(Bytes, len(e.bytes))
|
||||
|
||||
for i, c := range e.bytes {
|
||||
switch {
|
||||
case c >= 'A' && c <= 'Z':
|
||||
out[i] = 'A' + (c-'A'+13)%26
|
||||
case c >= 'a' && c <= 'z':
|
||||
out[i] = 'a' + (c-'a'+13)%26
|
||||
default:
|
||||
out[i] = c
|
||||
}
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
// Rot13 decodes the wrapped Bytes using the ROT13 cipher.
|
||||
func (d bdecode) Rot13() Bytes { return d.bytes.Encode().Rot13() }
|
||||
|
||||
// Octal returns the octal representation of the wrapped Bytes.
|
||||
//
|
||||
// Unlike String.Encode().Octal, which encodes Unicode code points, this
|
||||
// implementation is byte-wise: each byte is rendered as its octal value
|
||||
// (0-377), separated by spaces. The two representations match only for
|
||||
// ASCII input.
|
||||
func (e bencode) Octal() Bytes {
|
||||
var tmp [3]byte
|
||||
|
||||
out := make(Bytes, 0, len(e.bytes)*4)
|
||||
|
||||
for i, c := range e.bytes {
|
||||
if i != 0 {
|
||||
out = append(out, ' ')
|
||||
}
|
||||
|
||||
out = append(out, strconv.AppendUint(tmp[:0], uint64(c), 8)...)
|
||||
}
|
||||
|
||||
return out
|
||||
}
|
||||
|
||||
// Octal decodes the octal representation back to Bytes.
|
||||
// An empty input returns empty Bytes, mirroring bencode.Octal on empty input.
|
||||
// Each space-separated token must represent a valid byte value in the octal
|
||||
// range [0, 377]; anything else yields an error.
|
||||
func (d bdecode) Octal() Result[Bytes] {
|
||||
if d.bytes.IsEmpty() {
|
||||
return Ok(Bytes(""))
|
||||
}
|
||||
|
||||
out := make(Bytes, 0, (len(d.bytes)+1)/2)
|
||||
|
||||
for _, v := range d.bytes.Split(Bytes(" ")) {
|
||||
n, err := strconv.ParseUint(v.StringUnsafe().Std(), 8, 8)
|
||||
if err != nil {
|
||||
return Err[Bytes](err)
|
||||
}
|
||||
|
||||
out = append(out, byte(n))
|
||||
}
|
||||
|
||||
return Ok(out)
|
||||
}
|
||||
+521
@@ -0,0 +1,521 @@
|
||||
package g
|
||||
|
||||
import "strings"
|
||||
|
||||
// IsASCII checks if all bytes in the Bytes are ASCII bytes.
|
||||
func (bs Bytes) IsASCII() bool {
|
||||
for i := range bs {
|
||||
if bs[i] >= 0x80 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// IsDigit checks if the Bytes is non-empty and all bytes are ASCII digits ('0'-'9').
|
||||
// Unlike String.IsDigit, which is rune-aware and accepts Unicode digits, this method
|
||||
// operates byte-wise and only recognizes ASCII digits.
|
||||
func (bs Bytes) IsDigit() bool {
|
||||
if bs.IsEmpty() {
|
||||
return false
|
||||
}
|
||||
|
||||
for i := range bs {
|
||||
if bs[i] < '0' || bs[i] > '9' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// ReplaceMulti performs multiple replacements within the Bytes.
|
||||
//
|
||||
// The replacements are provided as pairs of old and new Bytes, in the same
|
||||
// order as in String.ReplaceMulti. Replacements are performed in a single pass:
|
||||
// at each position the first matching pattern in argument order wins, and
|
||||
// matches do not overlap. The number of arguments must be even; otherwise the
|
||||
// method panics.
|
||||
//
|
||||
// Both this method and String.ReplaceMulti match patterns as raw byte
|
||||
// sequences, so the two versions behave identically on the same data.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - oldnew ...Bytes: Pairs of Bytes to be replaced. Specify as many pairs as needed.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - Bytes: A new Bytes with replacements applied. The receiver is not modified.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// original := g.Bytes("Hello, world! This is a test.")
|
||||
// replaced := original.ReplaceMulti(
|
||||
// g.Bytes("Hello"), g.Bytes("Greetings"),
|
||||
// g.Bytes("world"), g.Bytes("universe"),
|
||||
// g.Bytes("test"), g.Bytes("example"),
|
||||
// )
|
||||
// // replaced contains "Greetings, universe! This is a example."
|
||||
func (bs Bytes) ReplaceMulti(oldnew ...Bytes) Bytes {
|
||||
pairs := make([]string, len(oldnew))
|
||||
for i, b := range oldnew {
|
||||
pairs[i] = string(b)
|
||||
}
|
||||
|
||||
return Bytes(strings.NewReplacer(pairs...).Replace(bs.StringUnsafe().Std()))
|
||||
}
|
||||
|
||||
// Remove removes all occurrences of the specified patterns from the Bytes.
|
||||
//
|
||||
// Both this method and String.Remove match patterns as raw byte sequences,
|
||||
// so the two versions behave identically on the same data.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - patterns ...Bytes: Patterns to be removed from the Bytes. Specify as many patterns as needed.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - Bytes: A new Bytes with all specified patterns removed. The receiver is
|
||||
// not modified. If no patterns are given, the original Bytes is returned.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// original := g.Bytes("Hello, world! This is a test.")
|
||||
// modified := original.Remove(
|
||||
// g.Bytes("Hello"),
|
||||
// g.Bytes("test"),
|
||||
// )
|
||||
// // modified contains ", world! This is a ."
|
||||
func (bs Bytes) Remove(patterns ...Bytes) Bytes {
|
||||
if len(patterns) == 0 {
|
||||
return bs
|
||||
}
|
||||
|
||||
pairs := make([]string, len(patterns)*2)
|
||||
for i, pattern := range patterns {
|
||||
pairs[i*2] = string(pattern)
|
||||
pairs[i*2+1] = ""
|
||||
}
|
||||
|
||||
return Bytes(strings.NewReplacer(pairs...).Replace(bs.StringUnsafe().Std()))
|
||||
}
|
||||
|
||||
// ReplaceNth returns a new Bytes with the nth occurrence of oldB
|
||||
// replaced with newB. If there aren't enough occurrences of oldB, the
|
||||
// original Bytes is returned. If n is less than -1, the original Bytes
|
||||
// is also returned. If n is -1, the last occurrence of oldB is replaced with newB.
|
||||
//
|
||||
// Both this method and String.ReplaceNth match patterns as raw byte sequences,
|
||||
// so the two versions behave identically on the same data.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - Bytes: A new Bytes with the nth occurrence of oldB replaced with newB.
|
||||
// The receiver is not modified.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// bs := g.Bytes("The quick brown dog jumped over the lazy dog.")
|
||||
// result := bs.ReplaceNth(g.Bytes("dog"), g.Bytes("fox"), 2)
|
||||
// fmt.Println(result)
|
||||
//
|
||||
// Output: "The quick brown dog jumped over the lazy fox.".
|
||||
func (bs Bytes) ReplaceNth(oldB, newB Bytes, n Int) Bytes {
|
||||
if n < -1 || oldB.IsEmpty() {
|
||||
return bs
|
||||
}
|
||||
|
||||
count, i := Int(0), Int(0)
|
||||
|
||||
for {
|
||||
pos := bs[i:].Index(oldB)
|
||||
if pos == -1 {
|
||||
break
|
||||
}
|
||||
|
||||
pos += i
|
||||
count++
|
||||
|
||||
if count == n || (n == -1 && bs[pos+oldB.Len():].Index(oldB) == -1) {
|
||||
result := make(Bytes, 0, bs.Len()+newB.Len()-oldB.Len())
|
||||
result = append(result, bs[:pos]...)
|
||||
result = append(result, newB...)
|
||||
result = append(result, bs[pos+oldB.Len():]...)
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
i = pos + oldB.Len()
|
||||
}
|
||||
|
||||
return bs
|
||||
}
|
||||
|
||||
// Chunks splits the Bytes into chunks of the specified size.
|
||||
//
|
||||
// This function iterates through the Bytes, creating chunks of the specified size.
|
||||
// If size is less than or equal to 0 or the Bytes is empty, it returns nil.
|
||||
// If size is greater than or equal to the length of the Bytes, it returns the
|
||||
// original Bytes as the only chunk.
|
||||
//
|
||||
// Unlike String.Chunks, which counts runes, the chunk size is measured in
|
||||
// bytes, so multibyte UTF-8 sequences may be split across chunks.
|
||||
//
|
||||
// The returned chunks are subslices sharing memory with the original Bytes,
|
||||
// as with Split; clone them if independent copies are needed.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - size (Int): The size of the chunks to split the Bytes into.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - []Bytes: the chunks of the specified size.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// bs := g.Bytes("Hello, World!")
|
||||
// chunks := bs.Chunks(4)
|
||||
//
|
||||
// chunks contains {"Hell", "o, W", "orld", "!"}.
|
||||
func (bs Bytes) Chunks(size Int) []Bytes {
|
||||
if size.Lte(0) || bs.IsEmpty() {
|
||||
return nil
|
||||
}
|
||||
|
||||
n := size.Std()
|
||||
l := len(bs)
|
||||
|
||||
if n >= l {
|
||||
return []Bytes{bs}
|
||||
}
|
||||
|
||||
result := make([]Bytes, 0, (l+n-1)/n)
|
||||
for i := 0; i < l; i += n {
|
||||
result = append(result, bs[i:min(i+n, l)])
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func (bs Bytes) Cut(start, end Bytes, rmtags ...bool) (Bytes, Bytes) {
|
||||
if start.IsEmpty() || end.IsEmpty() {
|
||||
return bs, Bytes("")
|
||||
}
|
||||
|
||||
startIndex := bs.Index(start)
|
||||
if startIndex == -1 {
|
||||
return bs, Bytes("")
|
||||
}
|
||||
|
||||
startEnd := startIndex + start.Len()
|
||||
endIndex := bs[startEnd:].Index(end)
|
||||
if endIndex == -1 {
|
||||
return bs, Bytes("")
|
||||
}
|
||||
|
||||
cut := bs[startEnd : startEnd+endIndex]
|
||||
|
||||
if len(rmtags) == 0 || !rmtags[0] {
|
||||
return bs, cut
|
||||
}
|
||||
|
||||
tail := startEnd + endIndex + end.Len()
|
||||
|
||||
remainder := make(Bytes, 0, startIndex+(bs.Len()-tail))
|
||||
remainder = append(remainder, bs[:startIndex]...)
|
||||
remainder = append(remainder, bs[tail:]...)
|
||||
|
||||
return remainder, cut
|
||||
}
|
||||
|
||||
// SubBytes extracts a subrange from the Bytes starting at the 'start' index and ending before the 'end' index.
|
||||
// The function also supports an optional 'step' parameter to define the increment between indices in the result.
|
||||
// If 'start' or 'end' index is negative, they represent positions relative to the end of the Bytes:
|
||||
// - A negative 'start' index indicates the position from the end of the Bytes, moving backward.
|
||||
// - A negative 'end' index indicates the position from the end of the Bytes.
|
||||
// The function ensures that indices are adjusted to fall within the valid range of the Bytes' length.
|
||||
// Out-of-bounds indices are clamped to the Bytes' bounds instead of panicking;
|
||||
// if 'start' exceeds 'end' (for a positive step) the result is an empty Bytes.
|
||||
//
|
||||
// Unlike String.SubString, which indexes runes, all indices and the step are
|
||||
// measured in bytes, so a boundary that falls inside a multibyte UTF-8 sequence
|
||||
// splits the rune and the result may not be valid UTF-8. A negative step
|
||||
// reverses bytes, not runes.
|
||||
//
|
||||
// The result is a newly allocated Bytes; the receiver is not modified.
|
||||
func (bs Bytes) SubBytes(start, end Int, step ...Int) Bytes {
|
||||
n := bs.Len()
|
||||
|
||||
clamp := func(i Int) Int {
|
||||
if i < 0 {
|
||||
i += n
|
||||
}
|
||||
|
||||
if i < 0 {
|
||||
return 0
|
||||
}
|
||||
|
||||
if i > n {
|
||||
return n
|
||||
}
|
||||
|
||||
return i
|
||||
}
|
||||
|
||||
start, end = clamp(start), clamp(end)
|
||||
|
||||
st := Int(1)
|
||||
if len(step) > 0 {
|
||||
st = step[0]
|
||||
}
|
||||
|
||||
// For a negative step the iteration starts AT start and moves down,
|
||||
// so a start clamped to n must begin at the last element.
|
||||
if st < 0 && start == n {
|
||||
start--
|
||||
}
|
||||
|
||||
if st == 1 {
|
||||
if start >= end {
|
||||
return Bytes{}
|
||||
}
|
||||
|
||||
return Bytes(append([]byte(nil), bs[start:end]...))
|
||||
}
|
||||
|
||||
if (start >= end && st > 0) || (start <= end && st < 0) || st == 0 {
|
||||
return Bytes{}
|
||||
}
|
||||
|
||||
var out []byte
|
||||
|
||||
if st > 0 {
|
||||
for i := start; i < end; i += st {
|
||||
out = append(out, bs[i])
|
||||
}
|
||||
} else {
|
||||
for i := start; i > end; i += st {
|
||||
out = append(out, bs[i])
|
||||
}
|
||||
}
|
||||
|
||||
return Bytes(out)
|
||||
}
|
||||
|
||||
// Similarity calculates the similarity between two Bytes using the
|
||||
// Levenshtein distance algorithm and returns the similarity percentage as a Float.
|
||||
//
|
||||
// The function compares two Bytes using the Levenshtein distance,
|
||||
// which measures the difference between two sequences by counting the number
|
||||
// of single-byte edits required to change one sequence into the other.
|
||||
// The similarity is then calculated by normalizing the distance by the maximum
|
||||
// length of the two input Bytes.
|
||||
//
|
||||
// Unlike String.Similarity, which compares runes, this method operates byte-wise,
|
||||
// so multibyte UTF-8 sequences are compared byte by byte.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - obs (Bytes): The Bytes to compare with bs.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - Float: The similarity percentage between the two Bytes as a value between 0 and 100.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// b1 := g.Bytes("kitten")
|
||||
// b2 := g.Bytes("sitting")
|
||||
// similarity := b1.Similarity(b2) // 57.14285714285714
|
||||
func (bs Bytes) Similarity(obs Bytes) Float {
|
||||
if bs.Eq(obs) {
|
||||
return 100
|
||||
}
|
||||
|
||||
if bs.IsEmpty() || obs.IsEmpty() {
|
||||
return 0
|
||||
}
|
||||
|
||||
s1, s2 := bs, obs
|
||||
|
||||
n1, n2 := len(s1), len(s2)
|
||||
|
||||
if n1 > n2 {
|
||||
s1, s2, n1, n2 = s2, s1, n2, n1
|
||||
}
|
||||
|
||||
distance := make([]int, n1+1)
|
||||
|
||||
for i, b2 := range s2 {
|
||||
prev := i + 1
|
||||
|
||||
for j, b1 := range s1 {
|
||||
current := distance[j]
|
||||
if b2 != b1 {
|
||||
current = min(distance[j]+1, min(prev+1, distance[j+1]+1))
|
||||
}
|
||||
|
||||
distance[j], prev = prev, current
|
||||
}
|
||||
|
||||
distance[n1] = prev
|
||||
}
|
||||
|
||||
return Float(1-float64(distance[n1])/float64(max(n1, n2))) * 100
|
||||
}
|
||||
|
||||
// Truncate shortens the Bytes to the specified maximum length. If the Bytes exceeds the
|
||||
// specified length, it is truncated, and an ellipsis ("...") is appended to indicate the truncation.
|
||||
//
|
||||
// If the length of the Bytes is less than or equal to the specified maximum length, the
|
||||
// original Bytes is returned unchanged.
|
||||
//
|
||||
// Unlike String.Truncate, which is rune-aware, this method truncates based on the number
|
||||
// of bytes, so multibyte UTF-8 sequences may be split.
|
||||
//
|
||||
// Parameters:
|
||||
// - max: The maximum number of bytes allowed in the resulting Bytes.
|
||||
//
|
||||
// Returns:
|
||||
// - A new Bytes truncated to the specified maximum length with "..." appended
|
||||
// if truncation occurs. Otherwise, returns the original Bytes.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// bs := g.Bytes("Hello, World!")
|
||||
// result := bs.Truncate(5)
|
||||
// // result: "Hello..."
|
||||
//
|
||||
// bs2 := g.Bytes("Short")
|
||||
// result2 := bs2.Truncate(10)
|
||||
// // result2: "Short"
|
||||
func (bs Bytes) Truncate(max Int) Bytes {
|
||||
if max.IsNegative() {
|
||||
return bs
|
||||
}
|
||||
|
||||
if bs.Len() <= max {
|
||||
return bs
|
||||
}
|
||||
|
||||
return append(bs[:max:max], "..."...)
|
||||
}
|
||||
|
||||
// LeftJustify justifies the Bytes to the left by adding padding to the right, up to the
|
||||
// specified length. If the length of the Bytes is already greater than or equal to the specified
|
||||
// length, or the pad is empty, the original Bytes is returned.
|
||||
//
|
||||
// The padding Bytes is repeated as necessary to fill the remaining length.
|
||||
// The padding is added to the right of the Bytes.
|
||||
//
|
||||
// Unlike String.LeftJustify, which counts runes, both length and padding are
|
||||
// measured in bytes.
|
||||
//
|
||||
// Parameters:
|
||||
// - length: The desired length of the resulting justified Bytes.
|
||||
// - pad: The Bytes used as padding.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// bs := g.Bytes("Hello")
|
||||
// result := bs.LeftJustify(10, g.Bytes("..."))
|
||||
// // result: "Hello....."
|
||||
func (bs Bytes) LeftJustify(length Int, pad Bytes) Bytes {
|
||||
if bs.Len() >= length || pad.IsEmpty() {
|
||||
return bs
|
||||
}
|
||||
|
||||
buf := make(Bytes, 0, length)
|
||||
buf = append(buf, bs...)
|
||||
buf = appendPadding(buf, pad, length-bs.Len())
|
||||
|
||||
return buf
|
||||
}
|
||||
|
||||
// RightJustify justifies the Bytes to the right by adding padding to the left, up to the
|
||||
// specified length. If the length of the Bytes is already greater than or equal to the specified
|
||||
// length, or the pad is empty, the original Bytes is returned.
|
||||
//
|
||||
// The padding Bytes is repeated as necessary to fill the remaining length.
|
||||
// The padding is added to the left of the Bytes.
|
||||
//
|
||||
// Unlike String.RightJustify, which counts runes, both length and padding are
|
||||
// measured in bytes.
|
||||
//
|
||||
// Parameters:
|
||||
// - length: The desired length of the resulting justified Bytes.
|
||||
// - pad: The Bytes used as padding.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// bs := g.Bytes("Hello")
|
||||
// result := bs.RightJustify(10, g.Bytes("..."))
|
||||
// // result: ".....Hello"
|
||||
func (bs Bytes) RightJustify(length Int, pad Bytes) Bytes {
|
||||
if bs.Len() >= length || pad.IsEmpty() {
|
||||
return bs
|
||||
}
|
||||
|
||||
buf := make(Bytes, 0, length)
|
||||
buf = appendPadding(buf, pad, length-bs.Len())
|
||||
buf = append(buf, bs...)
|
||||
|
||||
return buf
|
||||
}
|
||||
|
||||
// Center justifies the Bytes by adding padding on both sides, up to the specified length.
|
||||
// If the length of the Bytes is already greater than or equal to the specified length, or the
|
||||
// pad is empty, the original Bytes is returned.
|
||||
//
|
||||
// The padding Bytes is repeated as necessary to evenly distribute the remaining length on both
|
||||
// sides.
|
||||
// The padding is added to the left and right of the Bytes.
|
||||
//
|
||||
// Unlike String.Center, which counts runes, both length and padding are
|
||||
// measured in bytes.
|
||||
//
|
||||
// Parameters:
|
||||
// - length: The desired length of the resulting justified Bytes.
|
||||
// - pad: The Bytes used as padding.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// bs := g.Bytes("Hello")
|
||||
// result := bs.Center(10, g.Bytes("..."))
|
||||
// // result: "..Hello..."
|
||||
func (bs Bytes) Center(length Int, pad Bytes) Bytes {
|
||||
slen := bs.Len()
|
||||
if slen >= length || pad.IsEmpty() {
|
||||
return bs
|
||||
}
|
||||
|
||||
remains := length - slen
|
||||
|
||||
buf := make(Bytes, 0, length)
|
||||
buf = appendPadding(buf, pad, remains/2)
|
||||
buf = append(buf, bs...)
|
||||
buf = appendPadding(buf, pad, (remains+1)/2)
|
||||
|
||||
return buf
|
||||
}
|
||||
|
||||
// appendPadding appends the padding Bytes to buf to fill the remaining length.
|
||||
// It repeats the padding Bytes as necessary and appends any remaining bytes from
|
||||
// the padding Bytes.
|
||||
func appendPadding(buf, pad Bytes, remains Int) Bytes {
|
||||
padlen := pad.Len()
|
||||
|
||||
for range remains / padlen {
|
||||
buf = append(buf, pad...)
|
||||
}
|
||||
|
||||
if rem := remains % padlen; rem != 0 {
|
||||
buf = append(buf, pad[:rem]...)
|
||||
}
|
||||
|
||||
return buf
|
||||
}
|
||||
+77
@@ -0,0 +1,77 @@
|
||||
// Package g provides an ergonomic standard library extension for Go:
|
||||
// monadic error handling, rich generic containers, lazy iterators with
|
||||
// type-changing generic methods (Go 1.27+), and ergonomic wrappers for
|
||||
// primitive types and the filesystem.
|
||||
//
|
||||
// # Monads
|
||||
//
|
||||
// [Option] represents an optional value (Some/None) and [Result] represents
|
||||
// success or failure (Ok/Err). Both offer chainable combinators — Map, Then,
|
||||
// UnwrapOr, MapOr, Inspect and friends — so error and nil handling become
|
||||
// expressions instead of if-ladders.
|
||||
//
|
||||
// # Containers
|
||||
//
|
||||
// - [Slice]: an extended slice with 90+ methods.
|
||||
// - [Map]: a map with ergonomic accessors and entry API.
|
||||
// - [MapOrd]: an insertion-ordered map.
|
||||
// - [MapSafe]: a concurrency-safe map.
|
||||
// - [Set]: a hash set with algebraic operations (union, intersection, ...).
|
||||
// - [Deque]: a double-ended queue backed by a ring buffer.
|
||||
// - [Heap]: a binary min/max heap driven by a comparison function.
|
||||
//
|
||||
// # Iterators
|
||||
//
|
||||
// Every container exposes an Iter method returning lazy sequences (Heap
|
||||
// additionally offers a draining IntoIter) —
|
||||
// [Seq] for values, [Seq2] for key-value pairs, [SeqSlices] for grouped
|
||||
// windows/chunks and [SeqResult] for fallible streams — built on g's own
|
||||
// dependency-free iterator core. With Go 1.27 generic methods, transformations
|
||||
// can change the element type mid-chain:
|
||||
//
|
||||
// g.SliceOf(1, 2, 3).Iter().Map[string](strconv.Itoa).Collect().Slice() // Slice[string]
|
||||
//
|
||||
// Iterators are lazy: nothing is computed until a consumer (Collect, ForEach,
|
||||
// Fold, ...) runs the chain.
|
||||
//
|
||||
// # Primitive wrappers
|
||||
//
|
||||
// [String], [Int], [Float] and [Bytes] wrap the built-in types with fluent
|
||||
// methods, including conversion pipelines via Encode/Decode (Base64, Hex,
|
||||
// Octal, Binary, JSON, ...), Compress/Decompress (gzip, zlib, flate)
|
||||
// and Hash (MD5, SHA1, SHA256, SHA512).
|
||||
//
|
||||
// # Subpackages
|
||||
//
|
||||
// - fs: File and Dir — chainable, Result-based file and directory
|
||||
// operations, including lazy line/chunk iterators over file contents.
|
||||
// - rx: compiled regular expressions with a compile-once, rich matching API.
|
||||
// - rand: all randomness (numbers, choices, samples, shuffles, strings)
|
||||
// in one place.
|
||||
// - pool: a generic goroutine pool with limits, rate limiting and streaming.
|
||||
// - cmp: ordering primitives (cmp.Ordering, cmp.Cmp) used by sorts and heaps.
|
||||
// - f: predicate combinators for filters (f.Eq, f.Gt, f.Contains, ...).
|
||||
// - ref: pointer helpers (ref.Of).
|
||||
// - constraints: generic type constraints shared across the library.
|
||||
// - dbg: debugging helpers that print expressions with source locations.
|
||||
//
|
||||
// # Panics
|
||||
//
|
||||
// The library reports failures through [Result] and [Option]; panics are
|
||||
// reserved for programmer errors. Only two families of API panic:
|
||||
//
|
||||
// - the Unwrap/Expect family on [Option] and [Result] (Unwrap, UnwrapErr,
|
||||
// Expect) when called on the wrong variant;
|
||||
// - documented index-based operations (e.g. Slice.Swap/Insert/Replace/SubSlice,
|
||||
// Deque.Insert/Swap) when given out-of-range indices, and constructors with
|
||||
// documented preconditions (e.g. NewHeap with a nil comparison function).
|
||||
//
|
||||
// Everything else returns Option/Result instead of panicking.
|
||||
//
|
||||
// # Security
|
||||
//
|
||||
// [Format] placeholders resolve data only: each dot-segment selects a map
|
||||
// key, a MapOrd key, a slice/array index or a struct field — placeholders
|
||||
// cannot invoke methods. Note that an untrusted template string can still
|
||||
// choose which of the supplied argument data is printed.
|
||||
package g
|
||||
+178
@@ -0,0 +1,178 @@
|
||||
package g
|
||||
|
||||
// Entry is a sealed interface representing a view into a single Map entry.
|
||||
//
|
||||
// Entry provides an API for in-place manipulation of map entries, enabling
|
||||
// efficient "get or insert" patterns without redundant lookups.
|
||||
//
|
||||
// The interface is sealed to ensure type safety; implementations are limited
|
||||
// to [OccupiedEntry] (when the key exists) and [VacantEntry] (when the key
|
||||
// is absent). Use a type switch to access type-specific methods like Get,
|
||||
// Insert, or Remove.
|
||||
//
|
||||
// Common usage patterns:
|
||||
//
|
||||
// // Increment existing value or insert default
|
||||
// m.Entry("counter").AndModify(func(v *int) { *v++ }).OrInsert(1)
|
||||
//
|
||||
// // Insert only if absent
|
||||
// m.Entry("key").OrInsert(defaultValue)
|
||||
//
|
||||
// // Insert with lazy initialization
|
||||
// m.Entry("key").OrInsertWith(func() V { return expensiveComputation() })
|
||||
//
|
||||
// // Type switch for fine-grained control
|
||||
// switch e := m.Entry("key").(type) {
|
||||
// case OccupiedEntry[string, int]:
|
||||
// fmt.Println("exists:", e.Get())
|
||||
// case VacantEntry[string, int]:
|
||||
// e.Insert(42)
|
||||
// }
|
||||
//
|
||||
// An Entry is a short-lived view of the key state observed by Map.Entry.
|
||||
// Do not retain it across external insertion or removal of the same key;
|
||||
// obtain a fresh Entry after structurally changing that key.
|
||||
type Entry[K comparable, V any] interface {
|
||||
sealed()
|
||||
Key() K
|
||||
OrInsert(value V) V
|
||||
OrInsertWith(fn func() V) V
|
||||
OrInsertWithKey(fn func(K) V) V
|
||||
OrDefault() V
|
||||
AndModify(fn func(*V)) Entry[K, V]
|
||||
}
|
||||
|
||||
// OccupiedEntry represents a view into a map entry that is known to be present.
|
||||
//
|
||||
// It is typically obtained from Map.Entry(key) when the key already exists.
|
||||
// OccupiedEntry allows inspecting, modifying, replacing, or removing the value
|
||||
// associated with the key without performing additional map lookups.
|
||||
type OccupiedEntry[K comparable, V any] struct {
|
||||
m Map[K, V]
|
||||
key K
|
||||
}
|
||||
|
||||
// sealed prevents external implementations of the Entry interface.
|
||||
func (OccupiedEntry[K, V]) sealed() {}
|
||||
|
||||
// Key returns the key of this occupied entry.
|
||||
func (e OccupiedEntry[K, V]) Key() K { return e.key }
|
||||
|
||||
// Get returns the current value associated with the key.
|
||||
//
|
||||
// The value is returned by copy, consistent with Go map semantics.
|
||||
func (e OccupiedEntry[K, V]) Get() V { return e.m[e.key] }
|
||||
|
||||
// Insert replaces the value in the map with the provided one
|
||||
// and returns the previous value.
|
||||
//
|
||||
// The key remains present in the map.
|
||||
func (e OccupiedEntry[K, V]) Insert(value V) V {
|
||||
old := e.m[e.key]
|
||||
e.m[e.key] = value
|
||||
return old
|
||||
}
|
||||
|
||||
// Remove removes the entry from the map and returns the previously stored value.
|
||||
//
|
||||
// After this call, the key is no longer present in the map.
|
||||
func (e OccupiedEntry[K, V]) Remove() V {
|
||||
v := e.m[e.key]
|
||||
delete(e.m, e.key)
|
||||
return v
|
||||
}
|
||||
|
||||
// OrInsert returns the existing value without modifying the map.
|
||||
//
|
||||
// For OccupiedEntry, this is equivalent to Get since the key already exists.
|
||||
func (e OccupiedEntry[K, V]) OrInsert(value V) V { return e.Get() }
|
||||
|
||||
// OrInsertWith returns the existing value without invoking the function.
|
||||
//
|
||||
// For OccupiedEntry, the function is never called since the key already exists.
|
||||
func (e OccupiedEntry[K, V]) OrInsertWith(fn func() V) V { return e.Get() }
|
||||
|
||||
// OrInsertWithKey returns the existing value without invoking the function.
|
||||
//
|
||||
// For OccupiedEntry, the function is never called since the key already exists.
|
||||
func (e OccupiedEntry[K, V]) OrInsertWithKey(fn func(K) V) V { return e.Get() }
|
||||
|
||||
// OrDefault returns the existing value.
|
||||
//
|
||||
// For OccupiedEntry, this is equivalent to Get since the key already exists.
|
||||
func (e OccupiedEntry[K, V]) OrDefault() V { return e.Get() }
|
||||
|
||||
// AndModify applies the provided function to the value stored in the map
|
||||
// and returns the entry for method chaining.
|
||||
//
|
||||
// The function receives a pointer to a copy of the value; after modification,
|
||||
// the updated value is written back to the map.
|
||||
// The entry must not be used after the same key is externally removed or replaced.
|
||||
//
|
||||
// Example:
|
||||
//
|
||||
// m.Entry("count").AndModify(func(v *int) { *v++ }).OrInsert(1)
|
||||
func (e OccupiedEntry[K, V]) AndModify(fn func(*V)) Entry[K, V] {
|
||||
v := e.m[e.key]
|
||||
fn(&v)
|
||||
e.m[e.key] = v
|
||||
return e
|
||||
}
|
||||
|
||||
// VacantEntry represents a view into a map entry that is known to be absent.
|
||||
//
|
||||
// It is typically obtained from Map.Entry(key) when the key does not exist.
|
||||
// VacantEntry allows inserting a value for the key in a controlled manner.
|
||||
type VacantEntry[K comparable, V any] struct {
|
||||
m Map[K, V]
|
||||
key K
|
||||
}
|
||||
|
||||
// sealed prevents external implementations of the Entry interface.
|
||||
func (VacantEntry[K, V]) sealed() {}
|
||||
|
||||
// Key returns the key that would be used for insertion.
|
||||
func (e VacantEntry[K, V]) Key() K { return e.key }
|
||||
|
||||
// Insert inserts the provided value into the map and returns it.
|
||||
//
|
||||
// After this call, the key is present in the map with the given value.
|
||||
func (e VacantEntry[K, V]) Insert(value V) V {
|
||||
e.m[e.key] = value
|
||||
return value
|
||||
}
|
||||
|
||||
// OrInsert inserts the provided value and returns it.
|
||||
//
|
||||
// This is the primary method for inserting values via VacantEntry.
|
||||
func (e VacantEntry[K, V]) OrInsert(value V) V { return e.Insert(value) }
|
||||
|
||||
// OrInsertWith inserts the value returned by the function and returns it.
|
||||
//
|
||||
// The function is guaranteed to be called exactly once.
|
||||
// Use this when computing the default value is expensive.
|
||||
func (e VacantEntry[K, V]) OrInsertWith(fn func() V) V { return e.Insert(fn()) }
|
||||
|
||||
// OrInsertWithKey inserts the value returned by the function and returns it.
|
||||
//
|
||||
// The function receives the entry key and is guaranteed to be called exactly once.
|
||||
// Use this when the default value depends on the key.
|
||||
func (e VacantEntry[K, V]) OrInsertWithKey(fn func(K) V) V {
|
||||
return e.Insert(fn(e.key))
|
||||
}
|
||||
|
||||
// OrDefault inserts the zero value of V into the map and returns it.
|
||||
//
|
||||
// This is useful for types where the zero value is a valid initial state,
|
||||
// such as numeric types (0), slices (nil), or structs with zero defaults.
|
||||
func (e VacantEntry[K, V]) OrDefault() V {
|
||||
var zero V
|
||||
return e.Insert(zero)
|
||||
}
|
||||
|
||||
// AndModify does nothing for VacantEntry and returns the entry unchanged.
|
||||
//
|
||||
// Since there is no existing value to modify, the function is not called.
|
||||
// This allows fluent chaining like Entry(k).AndModify(f).OrInsert(v)
|
||||
// to work correctly regardless of whether the key exists.
|
||||
func (e VacantEntry[K, V]) AndModify(fn func(*V)) Entry[K, V] { return e }
|
||||
+184
@@ -0,0 +1,184 @@
|
||||
package g
|
||||
|
||||
import "slices"
|
||||
|
||||
// OrdEntry is a sealed interface representing a view into a single MapOrd entry.
|
||||
//
|
||||
// OrdEntry provides an API for in-place manipulation of ordered map entries,
|
||||
// enabling efficient "get or insert" patterns without redundant lookups while
|
||||
// preserving insertion order.
|
||||
//
|
||||
// The interface is sealed to ensure type safety; implementations are limited
|
||||
// to [OccupiedOrdEntry] (when the key exists) and [VacantOrdEntry] (when the
|
||||
// key is absent). Use a type switch to access type-specific methods like Get,
|
||||
// Insert, or Remove.
|
||||
//
|
||||
// Common usage patterns:
|
||||
//
|
||||
// // Increment existing value or insert default
|
||||
// mo.Entry("counter").AndModify(func(v *int) { *v++ }).OrInsert(1)
|
||||
//
|
||||
// // Insert only if absent (appends to end)
|
||||
// mo.Entry("key").OrInsert(defaultValue)
|
||||
//
|
||||
// // Type switch for fine-grained control
|
||||
// switch e := mo.Entry("key").(type) {
|
||||
// case OccupiedOrdEntry[string, int]:
|
||||
// fmt.Println("exists:", e.Get())
|
||||
// case VacantOrdEntry[string, int]:
|
||||
// e.Insert(42)
|
||||
// }
|
||||
type OrdEntry[K comparable, V any] interface {
|
||||
sealed()
|
||||
Key() K
|
||||
OrInsert(value V) V
|
||||
OrInsertWith(fn func() V) V
|
||||
OrInsertWithKey(fn func(K) V) V
|
||||
OrDefault() V
|
||||
AndModify(fn func(*V)) OrdEntry[K, V]
|
||||
}
|
||||
|
||||
// OccupiedOrdEntry represents a view into an ordered map entry that is known
|
||||
// to be present.
|
||||
//
|
||||
// It is typically obtained from MapOrd.Entry(key) when the key already exists.
|
||||
// OccupiedOrdEntry provides access to the key and the value stored in the
|
||||
// underlying ordered slice, allowing inspection, modification, replacement, or
|
||||
// removal. The key's position is resolved once, when the entry is created,
|
||||
// and reused directly by every operation. The entry is therefore only valid
|
||||
// as long as the MapOrd is not structurally mutated (Remove, SortBy, Clear)
|
||||
// between creation and use; obtain a fresh entry after such mutations.
|
||||
type OccupiedOrdEntry[K comparable, V any] struct {
|
||||
mo *MapOrd[K, V]
|
||||
key K
|
||||
idx int
|
||||
}
|
||||
|
||||
// sealed prevents external implementations of the OrdEntry interface.
|
||||
func (OccupiedOrdEntry[K, V]) sealed() {}
|
||||
|
||||
// Key returns the key of this occupied entry.
|
||||
func (e OccupiedOrdEntry[K, V]) Key() K { return e.key }
|
||||
|
||||
// Get returns the current value associated with the key.
|
||||
//
|
||||
// The value is returned by copy.
|
||||
func (e OccupiedOrdEntry[K, V]) Get() V { return (*e.mo)[e.idx].Value }
|
||||
|
||||
// Insert replaces the value at the entry's position with the provided value
|
||||
// and returns the previously stored value.
|
||||
//
|
||||
// The position of the entry in the ordered map is preserved.
|
||||
func (e OccupiedOrdEntry[K, V]) Insert(value V) V {
|
||||
old := (*e.mo)[e.idx].Value
|
||||
(*e.mo)[e.idx].Value = value
|
||||
|
||||
return old
|
||||
}
|
||||
|
||||
// Remove removes the entry from the ordered map and returns the previously
|
||||
// stored value.
|
||||
//
|
||||
// This operation preserves the relative order of the remaining entries.
|
||||
// After this call, the key is no longer present in the map and the entry
|
||||
// must not be used again.
|
||||
func (e OccupiedOrdEntry[K, V]) Remove() V {
|
||||
v := (*e.mo)[e.idx].Value
|
||||
*e.mo = slices.Delete(*e.mo, e.idx, e.idx+1)
|
||||
|
||||
return v
|
||||
}
|
||||
|
||||
// OrInsert returns the existing value without modifying the map.
|
||||
//
|
||||
// For OccupiedOrdEntry, this is equivalent to Get since the key already exists.
|
||||
func (e OccupiedOrdEntry[K, V]) OrInsert(value V) V { return e.Get() }
|
||||
|
||||
// OrInsertWith returns the existing value without invoking the function.
|
||||
//
|
||||
// For OccupiedOrdEntry, the function is never called since the key already exists.
|
||||
func (e OccupiedOrdEntry[K, V]) OrInsertWith(fn func() V) V { return e.Get() }
|
||||
|
||||
// OrInsertWithKey returns the existing value without invoking the function.
|
||||
//
|
||||
// For OccupiedOrdEntry, the function is never called since the key already exists.
|
||||
func (e OccupiedOrdEntry[K, V]) OrInsertWithKey(fn func(K) V) V { return e.Get() }
|
||||
|
||||
// OrDefault returns the existing value.
|
||||
//
|
||||
// For OccupiedOrdEntry, this is equivalent to Get since the key already exists.
|
||||
func (e OccupiedOrdEntry[K, V]) OrDefault() V { return e.Get() }
|
||||
|
||||
// AndModify applies the provided function to the value stored at the entry's
|
||||
// position and returns the entry for method chaining.
|
||||
//
|
||||
// The function receives a pointer to the actual value stored in the ordered map,
|
||||
// allowing in-place modification.
|
||||
//
|
||||
// Example:
|
||||
//
|
||||
// m.Entry("count").AndModify(func(v *int) { *v++ }).OrInsert(1)
|
||||
func (e OccupiedOrdEntry[K, V]) AndModify(fn func(*V)) OrdEntry[K, V] {
|
||||
fn(&(*e.mo)[e.idx].Value)
|
||||
return e
|
||||
}
|
||||
|
||||
// VacantOrdEntry represents a view into an ordered map entry that is known
|
||||
// to be absent.
|
||||
//
|
||||
// It is typically obtained from MapOrd.Entry(key) when the key does not exist.
|
||||
// VacantOrdEntry allows inserting a new key-value pair into the ordered map.
|
||||
type VacantOrdEntry[K comparable, V any] struct {
|
||||
mo *MapOrd[K, V]
|
||||
key K
|
||||
}
|
||||
|
||||
// sealed prevents external implementations of the OrdEntry interface.
|
||||
func (VacantOrdEntry[K, V]) sealed() {}
|
||||
|
||||
// Key returns the key that would be used for insertion.
|
||||
func (e VacantOrdEntry[K, V]) Key() K { return e.key }
|
||||
|
||||
// Insert inserts a new key-value pair into the ordered map and returns the value.
|
||||
//
|
||||
// The new entry is appended to the end of the ordered map.
|
||||
// After this call, the key is present in the map with the given value.
|
||||
func (e VacantOrdEntry[K, V]) Insert(value V) V {
|
||||
*e.mo = append(*e.mo, Pair[K, V]{Key: e.key, Value: value})
|
||||
return value
|
||||
}
|
||||
|
||||
// OrInsert inserts the provided value and returns it.
|
||||
//
|
||||
// This is the primary method for inserting values via VacantOrdEntry.
|
||||
func (e VacantOrdEntry[K, V]) OrInsert(value V) V { return e.Insert(value) }
|
||||
|
||||
// OrInsertWith inserts the value returned by the function and returns it.
|
||||
//
|
||||
// The function is guaranteed to be called exactly once.
|
||||
// Use this when computing the default value is expensive.
|
||||
func (e VacantOrdEntry[K, V]) OrInsertWith(fn func() V) V { return e.Insert(fn()) }
|
||||
|
||||
// OrInsertWithKey inserts the value returned by the function and returns it.
|
||||
//
|
||||
// The function receives the entry key and is guaranteed to be called exactly once.
|
||||
// Use this when the default value depends on the key.
|
||||
func (e VacantOrdEntry[K, V]) OrInsertWithKey(fn func(K) V) V {
|
||||
return e.Insert(fn(e.key))
|
||||
}
|
||||
|
||||
// OrDefault inserts the zero value of V into the ordered map and returns it.
|
||||
//
|
||||
// This is useful for types where the zero value is a valid initial state,
|
||||
// such as numeric types (0), slices (nil), or structs with zero defaults.
|
||||
func (e VacantOrdEntry[K, V]) OrDefault() V {
|
||||
var zero V
|
||||
return e.Insert(zero)
|
||||
}
|
||||
|
||||
// AndModify does nothing for VacantOrdEntry and returns the entry unchanged.
|
||||
//
|
||||
// Since there is no existing value to modify, the function is not called.
|
||||
// This allows fluent chaining like Entry(k).AndModify(f).OrInsert(v)
|
||||
// to work correctly regardless of whether the key exists.
|
||||
func (e VacantOrdEntry[K, V]) AndModify(fn func(*V)) OrdEntry[K, V] { return e }
|
||||
+335
@@ -0,0 +1,335 @@
|
||||
package g
|
||||
|
||||
// SafeEntry is a sealed interface representing a view into a single MapSafe entry.
|
||||
//
|
||||
// SafeEntry provides an API for in-place manipulation of concurrent map entries,
|
||||
// enabling efficient "get or insert" patterns that are safe for concurrent use
|
||||
// by multiple goroutines.
|
||||
//
|
||||
// The interface is sealed to ensure type safety; implementations are limited
|
||||
// to [OccupiedSafeEntry] (when the key exists) and [VacantSafeEntry] (when the
|
||||
// key is absent). Use a type switch to access type-specific methods like Get,
|
||||
// Insert, or Remove.
|
||||
//
|
||||
// Concurrency notes:
|
||||
// - AndModify uses a compare-and-swap (CAS) loop for atomic updates
|
||||
// - AndModify may invoke its callback more than once when CAS retries; the
|
||||
// callback must not perform non-idempotent external side effects
|
||||
// - VacantSafeEntry stores pending modifications to handle insertion races
|
||||
// - All operations are safe for concurrent use without external locking
|
||||
//
|
||||
// Common usage patterns:
|
||||
//
|
||||
// // Thread-safe increment or insert (safe for concurrent goroutines)
|
||||
// ms.Entry("counter").AndModify(func(v *int) { *v++ }).OrInsert(1)
|
||||
//
|
||||
// // Thread-safe insert only if absent
|
||||
// ms.Entry("key").OrInsert(defaultValue)
|
||||
//
|
||||
// // Type switch for fine-grained control
|
||||
// switch e := ms.Entry("key").(type) {
|
||||
// case OccupiedSafeEntry[string, int]:
|
||||
// fmt.Println("exists:", e.Get())
|
||||
// case VacantSafeEntry[string, int]:
|
||||
// e.Insert(42)
|
||||
// }
|
||||
type SafeEntry[K comparable, V any] interface {
|
||||
sealed()
|
||||
Key() K
|
||||
OrInsert(value V) V
|
||||
OrInsertWith(fn func() V) V
|
||||
OrInsertWithKey(fn func(K) V) V
|
||||
OrDefault() V
|
||||
AndModify(fn func(*V)) SafeEntry[K, V]
|
||||
}
|
||||
|
||||
// OccupiedSafeEntry represents a view into a concurrent map entry that is known
|
||||
// to be present.
|
||||
//
|
||||
// It is typically obtained from MapSafe.Entry(key) when the key exists.
|
||||
// All operations on OccupiedSafeEntry are safe for concurrent use.
|
||||
type OccupiedSafeEntry[K comparable, V any] struct {
|
||||
m *MapSafe[K, V]
|
||||
key K
|
||||
}
|
||||
|
||||
// sealed prevents external implementations of the SafeEntry interface.
|
||||
func (OccupiedSafeEntry[K, V]) sealed() {}
|
||||
|
||||
// Key returns the key of this entry.
|
||||
func (e OccupiedSafeEntry[K, V]) Key() K { return e.key }
|
||||
|
||||
// Get returns the current value associated with the key.
|
||||
//
|
||||
// If the key is concurrently removed, the zero value of V is returned.
|
||||
func (e OccupiedSafeEntry[K, V]) Get() V {
|
||||
if actual, ok := e.m.data.Load(e.key); ok {
|
||||
return *actual.(*V)
|
||||
}
|
||||
var zero V
|
||||
return zero
|
||||
}
|
||||
|
||||
// Insert replaces the value associated with the key and returns the previous value.
|
||||
//
|
||||
// The replacement is performed atomically with respect to other map operations.
|
||||
func (e OccupiedSafeEntry[K, V]) Insert(value V) V {
|
||||
e.m.structMu.RLock()
|
||||
defer e.m.structMu.RUnlock()
|
||||
|
||||
old, loaded := e.m.data.Swap(e.key, &value)
|
||||
if loaded {
|
||||
return *old.(*V)
|
||||
}
|
||||
|
||||
e.m.count.Add(1)
|
||||
var zero V
|
||||
return zero
|
||||
}
|
||||
|
||||
// Remove removes the entry from the map and returns the previously stored value.
|
||||
//
|
||||
// If the key is concurrently removed, the zero value of V is returned.
|
||||
func (e OccupiedSafeEntry[K, V]) Remove() V {
|
||||
e.m.structMu.RLock()
|
||||
defer e.m.structMu.RUnlock()
|
||||
|
||||
if actual, loaded := e.m.data.LoadAndDelete(e.key); loaded {
|
||||
e.m.count.Add(-1)
|
||||
return *actual.(*V)
|
||||
}
|
||||
|
||||
var zero V
|
||||
return zero
|
||||
}
|
||||
|
||||
// OrInsert returns the existing value without modifying the map.
|
||||
// If the key was concurrently removed, the value is inserted as for a vacant entry.
|
||||
func (e OccupiedSafeEntry[K, V]) OrInsert(value V) V {
|
||||
if actual, ok := e.m.data.Load(e.key); ok {
|
||||
return *actual.(*V)
|
||||
}
|
||||
|
||||
return VacantSafeEntry[K, V]{m: e.m, key: e.key}.OrInsert(value)
|
||||
}
|
||||
|
||||
// OrInsertWith returns the existing value if present, or inserts the result
|
||||
// of fn() and returns it.
|
||||
//
|
||||
// Note: Due to concurrent access, fn() may be invoked even if another
|
||||
// goroutine inserts the key between the check and insertion. In this case,
|
||||
// the result of fn() is discarded and the existing value is returned.
|
||||
func (e OccupiedSafeEntry[K, V]) OrInsertWith(fn func() V) V {
|
||||
if actual, ok := e.m.data.Load(e.key); ok {
|
||||
return *actual.(*V)
|
||||
}
|
||||
|
||||
return VacantSafeEntry[K, V]{m: e.m, key: e.key}.OrInsertWith(fn)
|
||||
}
|
||||
|
||||
// OrInsertWithKey returns the existing value without invoking the function.
|
||||
// If the key was concurrently removed, fn is invoked with the key and its
|
||||
// result inserted as for a vacant entry.
|
||||
func (e OccupiedSafeEntry[K, V]) OrInsertWithKey(fn func(K) V) V {
|
||||
if actual, ok := e.m.data.Load(e.key); ok {
|
||||
return *actual.(*V)
|
||||
}
|
||||
|
||||
return VacantSafeEntry[K, V]{m: e.m, key: e.key}.OrInsertWithKey(fn)
|
||||
}
|
||||
|
||||
// OrDefault returns the existing value.
|
||||
// If the key was concurrently removed, the zero value of V is inserted and returned.
|
||||
func (e OccupiedSafeEntry[K, V]) OrDefault() V {
|
||||
if actual, ok := e.m.data.Load(e.key); ok {
|
||||
return *actual.(*V)
|
||||
}
|
||||
|
||||
var zero V
|
||||
return VacantSafeEntry[K, V]{m: e.m, key: e.key}.Insert(zero)
|
||||
}
|
||||
|
||||
// AndModify applies the provided function to the value associated with the key
|
||||
// and returns the entry.
|
||||
//
|
||||
// The modification is performed using a compare-and-swap loop.
|
||||
// The function receives a pointer to a copy of the value; the updated value
|
||||
// is written back atomically.
|
||||
// Under contention, fn may be invoked more than once before CAS succeeds.
|
||||
// It should only modify the provided value and must not rely on exactly-once
|
||||
// external side effects.
|
||||
//
|
||||
// If the key is concurrently removed, AndModify becomes a no-op.
|
||||
func (e OccupiedSafeEntry[K, V]) AndModify(fn func(*V)) SafeEntry[K, V] {
|
||||
for {
|
||||
actual, ok := e.m.data.Load(e.key)
|
||||
if !ok {
|
||||
return e
|
||||
}
|
||||
|
||||
oldPtr := actual.(*V)
|
||||
newVal := *oldPtr
|
||||
fn(&newVal)
|
||||
|
||||
if e.m.data.CompareAndSwap(e.key, oldPtr, &newVal) {
|
||||
return e
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// VacantSafeEntry represents a view into a concurrent map entry that is known
|
||||
// to be absent at the time of creation.
|
||||
//
|
||||
// It is typically obtained from MapSafe.Entry(key) when the key does not exist.
|
||||
// All operations on VacantSafeEntry are safe for concurrent use.
|
||||
//
|
||||
// The modify field stores a pending modification function from AndModify,
|
||||
// which will be applied if OrInsert loses a race with another goroutine.
|
||||
type VacantSafeEntry[K comparable, V any] struct {
|
||||
m *MapSafe[K, V]
|
||||
key K
|
||||
modify func(*V)
|
||||
}
|
||||
|
||||
// sealed prevents external implementations of the SafeEntry interface.
|
||||
func (VacantSafeEntry[K, V]) sealed() {}
|
||||
|
||||
// Key returns the key that would be used for insertion.
|
||||
func (e VacantSafeEntry[K, V]) Key() K { return e.key }
|
||||
|
||||
// applyPending applies the pending AndModify function (if any) to the value now
|
||||
// stored under the key, then returns the resulting stored value.
|
||||
//
|
||||
// It is used by the Or* methods when an insert loses the race to a concurrent
|
||||
// goroutine: the modification registered via AndModify must still be applied to
|
||||
// the value that actually won the race. If the key has since been removed,
|
||||
// fallback (the value observed during the failed insert) is returned instead.
|
||||
func (e VacantSafeEntry[K, V]) applyPending(fallback *V) V {
|
||||
if e.modify != nil {
|
||||
OccupiedSafeEntry[K, V]{m: e.m, key: e.key}.AndModify(e.modify)
|
||||
if val, ok := e.m.data.Load(e.key); ok {
|
||||
return *val.(*V)
|
||||
}
|
||||
}
|
||||
|
||||
return *fallback
|
||||
}
|
||||
|
||||
// Insert inserts the provided value into the map and returns the stored value.
|
||||
//
|
||||
// If another goroutine inserts the same key concurrently, the existing value
|
||||
// is returned instead. Equivalent to OrInsert for VacantSafeEntry.
|
||||
func (e VacantSafeEntry[K, V]) Insert(value V) V { return e.OrInsert(value) }
|
||||
|
||||
// OrInsert inserts the provided value and returns the stored value.
|
||||
//
|
||||
// If another goroutine inserts the same key concurrently (the insert "loses
|
||||
// the race"), the existing value is used instead. In this case, if a pending
|
||||
// modification was registered via AndModify, it is applied atomically to the
|
||||
// existing value before returning.
|
||||
//
|
||||
// This ensures that chained calls like Entry(k).AndModify(f).OrInsert(v)
|
||||
// behave correctly under concurrent access: under contended insertion the
|
||||
// modification is applied to the winning value.
|
||||
//
|
||||
// Edge case: if the key is also removed by another goroutine between the lost
|
||||
// insert and the modify, the modification has nothing to apply to and the
|
||||
// pre-modify value observed during the failed insert is returned. In that
|
||||
// narrow window the modification can be lost.
|
||||
func (e VacantSafeEntry[K, V]) OrInsert(value V) V {
|
||||
e.m.structMu.RLock()
|
||||
defer e.m.structMu.RUnlock()
|
||||
|
||||
actual, loaded := e.m.data.LoadOrStore(e.key, &value)
|
||||
if !loaded {
|
||||
e.m.count.Add(1)
|
||||
}
|
||||
|
||||
if loaded {
|
||||
return e.applyPending(actual.(*V))
|
||||
}
|
||||
|
||||
return *actual.(*V)
|
||||
}
|
||||
|
||||
// OrInsertWith inserts the value returned by the function and returns the
|
||||
// stored value.
|
||||
//
|
||||
// Note: Due to lock-free implementation, fn() is evaluated before the atomic
|
||||
// insertion. If another goroutine inserts the key concurrently, the result
|
||||
// of fn() may be discarded and the existing value is returned instead.
|
||||
func (e VacantSafeEntry[K, V]) OrInsertWith(fn func() V) V {
|
||||
if actual, ok := e.m.data.Load(e.key); ok {
|
||||
return e.applyPending(actual.(*V))
|
||||
}
|
||||
|
||||
e.m.structMu.RLock()
|
||||
defer e.m.structMu.RUnlock()
|
||||
|
||||
actual, loaded := e.m.data.LoadOrStore(e.key, new(fn()))
|
||||
if !loaded {
|
||||
e.m.count.Add(1)
|
||||
}
|
||||
|
||||
if loaded {
|
||||
return e.applyPending(actual.(*V))
|
||||
}
|
||||
|
||||
return *actual.(*V)
|
||||
}
|
||||
|
||||
// OrInsertWithKey inserts the value returned by the function and returns the
|
||||
// stored value.
|
||||
//
|
||||
// Note: Due to lock-free implementation, fn() is evaluated before the atomic
|
||||
// insertion. If another goroutine inserts the key concurrently, the result
|
||||
// of fn() may be discarded and the existing value is returned instead.
|
||||
func (e VacantSafeEntry[K, V]) OrInsertWithKey(fn func(K) V) V {
|
||||
if actual, ok := e.m.data.Load(e.key); ok {
|
||||
return e.applyPending(actual.(*V))
|
||||
}
|
||||
|
||||
e.m.structMu.RLock()
|
||||
defer e.m.structMu.RUnlock()
|
||||
|
||||
actual, loaded := e.m.data.LoadOrStore(e.key, new(fn(e.key)))
|
||||
if !loaded {
|
||||
e.m.count.Add(1)
|
||||
}
|
||||
|
||||
if loaded {
|
||||
return e.applyPending(actual.(*V))
|
||||
}
|
||||
|
||||
return *actual.(*V)
|
||||
}
|
||||
|
||||
// OrDefault inserts the zero value of V into the map and returns the stored value.
|
||||
//
|
||||
// If another goroutine inserts the same key concurrently and a pending
|
||||
// modification was registered via AndModify, it is applied atomically
|
||||
// to the existing value before returning.
|
||||
func (e VacantSafeEntry[K, V]) OrDefault() V {
|
||||
var zero V
|
||||
return e.OrInsert(zero)
|
||||
}
|
||||
|
||||
// AndModify registers a modification function to be applied to the value.
|
||||
//
|
||||
// If the key was concurrently inserted by another goroutine since this
|
||||
// VacantSafeEntry was created, the modification is applied immediately
|
||||
// via OccupiedSafeEntry.AndModify.
|
||||
//
|
||||
// Otherwise, the function is stored and will be applied later by OrInsert
|
||||
// if it loses the race to insert the key. This ensures that the pattern
|
||||
// Entry(k).AndModify(f).OrInsert(v) correctly increments existing values
|
||||
// even under heavy concurrent access.
|
||||
//
|
||||
// Returns the appropriate SafeEntry for method chaining.
|
||||
func (e VacantSafeEntry[K, V]) AndModify(fn func(*V)) SafeEntry[K, V] {
|
||||
if _, ok := e.m.data.Load(e.key); ok {
|
||||
return OccupiedSafeEntry[K, V]{m: e.m, key: e.key}.AndModify(fn)
|
||||
}
|
||||
|
||||
return VacantSafeEntry[K, V]{m: e.m, key: e.key, modify: fn}
|
||||
}
|
||||
+158
@@ -0,0 +1,158 @@
|
||||
package g
|
||||
|
||||
import "math"
|
||||
|
||||
// IsNaN reports whether the Float is an IEEE 754 "not-a-number" value.
|
||||
func (f Float) IsNaN() bool { return math.IsNaN(f.Std()) }
|
||||
|
||||
// IsInf reports whether the Float is an infinity, either positive or negative.
|
||||
func (f Float) IsInf() bool { return math.IsInf(f.Std(), 0) }
|
||||
|
||||
// IsFinite reports whether the Float is neither NaN nor an infinity.
|
||||
func (f Float) IsFinite() bool { return !f.IsNaN() && !f.IsInf() }
|
||||
|
||||
// IsNormal reports whether the Float is a normal IEEE 754 number:
|
||||
// neither zero, subnormal, infinite, nor NaN.
|
||||
func (f Float) IsNormal() bool {
|
||||
exp := f.Bits() >> 52 & 0x7ff
|
||||
|
||||
return exp != 0 && exp != 0x7ff
|
||||
}
|
||||
|
||||
// Signum returns a Float representing the sign of the Float:
|
||||
// 1 if the sign bit is clear (including +0), -1 if the sign bit is set (including -0),
|
||||
// and NaN if the Float is NaN.
|
||||
func (f Float) Signum() Float {
|
||||
if f.IsNaN() {
|
||||
return f
|
||||
}
|
||||
|
||||
return Float(math.Copysign(1, f.Std()))
|
||||
}
|
||||
|
||||
// IsSignPositive reports whether the Float has a positive sign bit.
|
||||
// This includes +0.0 and positive infinity.
|
||||
// Note: NaN carries a sign bit too, so a NaN with a clear sign bit (e.g. math.NaN())
|
||||
// is reported as sign-positive; use IsNaN to detect NaN itself.
|
||||
func (f Float) IsSignPositive() bool { return !math.Signbit(f.Std()) }
|
||||
|
||||
// IsSignNegative reports whether the Float has a negative sign bit.
|
||||
// This includes -0.0 and negative infinity.
|
||||
// Note: NaN carries a sign bit too, so a NaN with a set sign bit (e.g.
|
||||
// math.Copysign(math.NaN(), -1)) is reported as sign-negative; use IsNaN to
|
||||
// detect NaN itself.
|
||||
func (f Float) IsSignNegative() bool { return math.Signbit(f.Std()) }
|
||||
|
||||
// Ceil returns the least integer value greater than or equal to the Float.
|
||||
func (f Float) Ceil() Float { return Float(math.Ceil(f.Std())) }
|
||||
|
||||
// Floor returns the greatest integer value less than or equal to the Float.
|
||||
func (f Float) Floor() Float { return Float(math.Floor(f.Std())) }
|
||||
|
||||
// Trunc returns the integer part of the Float, rounding toward zero.
|
||||
func (f Float) Trunc() Float { return Float(math.Trunc(f.Std())) }
|
||||
|
||||
// Fract returns the fractional part of the Float (f - f.Trunc()).
|
||||
// For NaN and ±Inf the result is NaN.
|
||||
func (f Float) Fract() Float { return f - f.Trunc() }
|
||||
|
||||
// Clamp restricts the Float to the inclusive range [min, max].
|
||||
// If the Float is NaN, NaN is returned.
|
||||
// The caller must ensure min <= max and that neither bound is NaN: this method
|
||||
// does not panic on an invalid range — a NaN bound never
|
||||
// compares true, so the corresponding check is silently skipped, and with
|
||||
// min > max the lower bound wins.
|
||||
func (f Float) Clamp(min, max Float) Float {
|
||||
if f < min {
|
||||
return min
|
||||
}
|
||||
|
||||
if f > max {
|
||||
return max
|
||||
}
|
||||
|
||||
return f
|
||||
}
|
||||
|
||||
// Recip returns the reciprocal (multiplicative inverse) of the Float, 1/f.
|
||||
func (f Float) Recip() Float { return 1 / f }
|
||||
|
||||
// Copysign returns a Float with the magnitude of the Float and the sign of sign.
|
||||
func (f Float) Copysign(sign Float) Float { return Float(math.Copysign(f.Std(), sign.Std())) }
|
||||
|
||||
// MulAdd returns f*b + c computed as a fused multiply-add with only one rounding.
|
||||
func (f Float) MulAdd(b, c Float) Float { return Float(math.FMA(f.Std(), b.Std(), c.Std())) }
|
||||
|
||||
// Hypot returns Sqrt(f*f + b*b), avoiding unnecessary overflow and underflow.
|
||||
func (f Float) Hypot(b Float) Float { return Float(math.Hypot(f.Std(), b.Std())) }
|
||||
|
||||
// Cbrt returns the cube root of the Float.
|
||||
func (f Float) Cbrt() Float { return Float(math.Cbrt(f.Std())) }
|
||||
|
||||
// Exp returns e**f, the base-e exponential of the Float.
|
||||
func (f Float) Exp() Float { return Float(math.Exp(f.Std())) }
|
||||
|
||||
// Exp2 returns 2**f, the base-2 exponential of the Float.
|
||||
func (f Float) Exp2() Float { return Float(math.Exp2(f.Std())) }
|
||||
|
||||
// ExpM1 returns e**f - 1, which is more accurate than Exp().Sub(1) when the Float is near zero.
|
||||
func (f Float) ExpM1() Float { return Float(math.Expm1(f.Std())) }
|
||||
|
||||
// Ln returns the natural logarithm of the Float.
|
||||
func (f Float) Ln() Float { return Float(math.Log(f.Std())) }
|
||||
|
||||
// Ln1p returns the natural logarithm of 1 plus the Float,
|
||||
// which is more accurate than Add(1).Ln() when the Float is near zero.
|
||||
func (f Float) Ln1p() Float { return Float(math.Log1p(f.Std())) }
|
||||
|
||||
// Log2 returns the base-2 logarithm of the Float.
|
||||
func (f Float) Log2() Float { return Float(math.Log2(f.Std())) }
|
||||
|
||||
// Log10 returns the base-10 logarithm of the Float.
|
||||
func (f Float) Log10() Float { return Float(math.Log10(f.Std())) }
|
||||
|
||||
// Sin returns the sine of the Float (in radians).
|
||||
func (f Float) Sin() Float { return Float(math.Sin(f.Std())) }
|
||||
|
||||
// Cos returns the cosine of the Float (in radians).
|
||||
func (f Float) Cos() Float { return Float(math.Cos(f.Std())) }
|
||||
|
||||
// Tan returns the tangent of the Float (in radians).
|
||||
func (f Float) Tan() Float { return Float(math.Tan(f.Std())) }
|
||||
|
||||
// Asin returns the arcsine of the Float, in radians.
|
||||
func (f Float) Asin() Float { return Float(math.Asin(f.Std())) }
|
||||
|
||||
// Acos returns the arccosine of the Float, in radians.
|
||||
func (f Float) Acos() Float { return Float(math.Acos(f.Std())) }
|
||||
|
||||
// Atan returns the arctangent of the Float, in radians.
|
||||
func (f Float) Atan() Float { return Float(math.Atan(f.Std())) }
|
||||
|
||||
// Atan2 returns the arctangent of f/b (with f as y and b as x), in radians,
|
||||
// using the signs of the two to determine the quadrant of the result.
|
||||
func (f Float) Atan2(b Float) Float { return Float(math.Atan2(f.Std(), b.Std())) }
|
||||
|
||||
// Sinh returns the hyperbolic sine of the Float.
|
||||
func (f Float) Sinh() Float { return Float(math.Sinh(f.Std())) }
|
||||
|
||||
// Cosh returns the hyperbolic cosine of the Float.
|
||||
func (f Float) Cosh() Float { return Float(math.Cosh(f.Std())) }
|
||||
|
||||
// Tanh returns the hyperbolic tangent of the Float.
|
||||
func (f Float) Tanh() Float { return Float(math.Tanh(f.Std())) }
|
||||
|
||||
// Asinh returns the inverse hyperbolic sine of the Float.
|
||||
func (f Float) Asinh() Float { return Float(math.Asinh(f.Std())) }
|
||||
|
||||
// Acosh returns the inverse hyperbolic cosine of the Float.
|
||||
func (f Float) Acosh() Float { return Float(math.Acosh(f.Std())) }
|
||||
|
||||
// Atanh returns the inverse hyperbolic tangent of the Float.
|
||||
func (f Float) Atanh() Float { return Float(math.Atanh(f.Std())) }
|
||||
|
||||
// ToDegrees converts the Float from radians to degrees.
|
||||
func (f Float) ToDegrees() Float { return f * (180 / math.Pi) }
|
||||
|
||||
// ToRadians converts the Float from degrees to radians.
|
||||
func (f Float) ToRadians() Float { return f * (math.Pi / 180) }
|
||||
+580
@@ -0,0 +1,580 @@
|
||||
package fs
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/enetx/g"
|
||||
)
|
||||
|
||||
// Dir is a struct representing a directory path.
|
||||
type Dir struct {
|
||||
path g.String // Directory path.
|
||||
}
|
||||
|
||||
// NewDir returns a new Dir instance with the given path.
|
||||
func NewDir(path g.String) *Dir { return &Dir{path: path} }
|
||||
|
||||
// Chown changes the ownership of the directory to the specified UID and GID.
|
||||
// It uses os.Chown to modify ownership and returns a Result[*Dir] indicating success or failure.
|
||||
func (d *Dir) Chown(uid, gid int) g.Result[*Dir] {
|
||||
err := os.Chown(d.path.Std(), uid, gid)
|
||||
if err != nil {
|
||||
return g.Err[*Dir](err)
|
||||
}
|
||||
|
||||
return g.Ok(d)
|
||||
}
|
||||
|
||||
// Stat retrieves information about the directory represented by the Dir instance.
|
||||
// It returns a Result[fs.FileInfo] containing details about the directory's metadata.
|
||||
func (d *Dir) Stat() g.Result[fs.FileInfo] {
|
||||
path := d.Path()
|
||||
if path.IsErr() {
|
||||
return g.Err[fs.FileInfo](path.Err())
|
||||
}
|
||||
|
||||
return g.ResultOf(os.Stat(path.Ok().Std()))
|
||||
}
|
||||
|
||||
// Lstat retrieves information about the symbolic link represented by the Dir instance.
|
||||
// It returns a Result[fs.FileInfo] containing details about the symbolic link's metadata.
|
||||
// Unlike Stat, Lstat does not follow the link and provides information about the link itself.
|
||||
func (d *Dir) Lstat() g.Result[fs.FileInfo] {
|
||||
path := d.Path()
|
||||
if path.IsErr() {
|
||||
return g.Err[fs.FileInfo](path.Err())
|
||||
}
|
||||
|
||||
return g.ResultOf(os.Lstat(path.Ok().Std()))
|
||||
}
|
||||
|
||||
// IsLink checks if the directory is a symbolic link.
|
||||
func (d *Dir) IsLink() bool {
|
||||
stat := d.Lstat()
|
||||
return stat.IsOk() && stat.Ok().Mode()&os.ModeSymlink != 0
|
||||
}
|
||||
|
||||
// CreateTempDir creates a new temporary directory in the specified directory with the
|
||||
// specified name pattern and returns a Result, which contains a pointer to the Dir
|
||||
// or an error if the operation fails.
|
||||
// If no directory is specified, the default directory for temporary directories is used.
|
||||
// If no name pattern is specified, the default pattern "*" is used.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - args ...String: A variadic parameter specifying the directory and/or name
|
||||
// pattern for the temporary directory.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - *Dir: A pointer to the Dir representing the temporary directory.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// tmpdir := fs.CreateTempDir() // Creates a temporary directory with default settings
|
||||
// tmpdirWithDir := fs.CreateTempDir("mydir") // Creates a temporary directory in "mydir" directory
|
||||
// tmpdirWithPattern := fs.CreateTempDir("", "tmp") // Creates a temporary directory with "tmp" pattern
|
||||
func CreateTempDir(args ...g.String) g.Result[*Dir] {
|
||||
dir := ""
|
||||
pattern := "*"
|
||||
|
||||
if len(args) != 0 {
|
||||
if len(args) > 1 {
|
||||
pattern = args[1].Std()
|
||||
}
|
||||
|
||||
dir = args[0].Std()
|
||||
}
|
||||
|
||||
tmpDir, err := os.MkdirTemp(dir, pattern)
|
||||
if err != nil {
|
||||
return g.Err[*Dir](err)
|
||||
}
|
||||
|
||||
return g.Ok(NewDir(g.String(tmpDir)))
|
||||
}
|
||||
|
||||
// TempDir returns the default directory to use for temporary files.
|
||||
//
|
||||
// On Unix systems, it returns $TMPDIR if non-empty, else /tmp.
|
||||
// On Windows, it uses GetTempPath, returning the first non-empty
|
||||
// value from %TMP%, %TEMP%, %USERPROFILE%, or the Windows directory.
|
||||
// On Plan 9, it returns /tmp.
|
||||
//
|
||||
// The directory is neither guaranteed to exist nor have accessible
|
||||
// permissions.
|
||||
func TempDir() *Dir { return NewDir(g.String(os.TempDir())) }
|
||||
|
||||
// Remove attempts to delete the directory and its contents.
|
||||
// It returns a Result, which contains either the *Dir or an error.
|
||||
// If the directory does not exist, Remove returns a successful Result with *Dir set.
|
||||
// Any error that occurs during removal will be of type *PathError.
|
||||
func (d *Dir) Remove() g.Result[*Dir] {
|
||||
if err := os.RemoveAll(d.String().Std()); err != nil {
|
||||
return g.Err[*Dir](err)
|
||||
}
|
||||
|
||||
return g.Ok(d)
|
||||
}
|
||||
|
||||
// Copy copies the contents of the current directory to the destination directory.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - dest (String): The destination directory where the contents of the current directory should be copied.
|
||||
//
|
||||
// - followLinks (optional): A boolean indicating whether to follow symbolic links during the walk.
|
||||
// If true, symbolic links are followed; otherwise, they are skipped.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - Result[*Dir]: A Result type containing either a pointer to a new Dir instance representing the destination directory or an error.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// sourceDir := fs.NewDir("path/to/source")
|
||||
// destinationDirResult := sourceDir.Copy("path/to/destination")
|
||||
// if destinationDirResult.IsErr() {
|
||||
// // Handle error
|
||||
// }
|
||||
// destinationDir := destinationDirResult.Ok()
|
||||
func (d *Dir) Copy(dest g.String, followLinks ...bool) g.Result[*Dir] {
|
||||
files := g.NewSlice[*File]()
|
||||
|
||||
for r := range d.Walk() {
|
||||
if r.IsErr() {
|
||||
return g.Err[*Dir](r.Err())
|
||||
}
|
||||
files.Push(r.Ok())
|
||||
}
|
||||
|
||||
root := d.Path()
|
||||
if root.IsErr() {
|
||||
return g.Err[*Dir](root.Err())
|
||||
}
|
||||
|
||||
destRoot := NewDir(dest).Path()
|
||||
if destRoot.IsErr() {
|
||||
return g.Err[*Dir](destRoot.Err())
|
||||
}
|
||||
|
||||
follow := true
|
||||
if len(followLinks) > 0 {
|
||||
follow = followLinks[0]
|
||||
}
|
||||
|
||||
for f := range files.Iter() {
|
||||
path := f.Path()
|
||||
if path.IsErr() {
|
||||
return g.Err[*Dir](path.Err())
|
||||
}
|
||||
|
||||
relpath, err := filepath.Rel(root.Ok().Std(), path.Ok().Std())
|
||||
if err != nil {
|
||||
return g.Err[*Dir](err)
|
||||
}
|
||||
|
||||
destpath := g.String(filepath.Join(destRoot.Ok().Std(), relpath))
|
||||
|
||||
// Skip every symlink (file or directory) when not following links;
|
||||
// a symlink to a regular file is not an IsDir entry, so the skip must
|
||||
// happen before the IsDir branch to avoid dereferencing the link.
|
||||
if !follow && f.IsLink() {
|
||||
continue
|
||||
}
|
||||
|
||||
stat := f.Stat()
|
||||
if stat.IsErr() {
|
||||
return g.Err[*Dir](stat.Err())
|
||||
}
|
||||
|
||||
if stat.Ok().IsDir() {
|
||||
if r := NewDir(destpath).CreateAll(stat.Ok().Mode()); r.IsErr() {
|
||||
return r
|
||||
}
|
||||
|
||||
continue
|
||||
}
|
||||
|
||||
if r := f.Copy(destpath, stat.Ok().Mode()); r.IsErr() {
|
||||
return g.Err[*Dir](r.Err())
|
||||
}
|
||||
}
|
||||
|
||||
return g.Ok(NewDir(dest))
|
||||
}
|
||||
|
||||
// Create creates a new directory with the specified mode (optional).
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - mode (os.FileMode, optional): The file mode for the new directory.
|
||||
// If not provided, it defaults to DirDefault (0755).
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - *Dir: A pointer to the Dir instance on which the method was called.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// dir := fs.NewDir("path/to/directory")
|
||||
// createdDir := dir.Create(0755) // Optional mode argument
|
||||
func (d *Dir) Create(mode ...os.FileMode) g.Result[*Dir] {
|
||||
dmode := os.FileMode(g.DirDefault)
|
||||
if len(mode) > 0 {
|
||||
dmode = mode[0]
|
||||
}
|
||||
if err := os.Mkdir(d.path.Std(), dmode); err != nil {
|
||||
return g.Err[*Dir](err)
|
||||
}
|
||||
|
||||
return g.Ok(d)
|
||||
}
|
||||
|
||||
// Join joins the current directory path with the given path elements, returning the joined path.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - elem (...String): One or more String values representing path elements to
|
||||
// be joined with the current directory path.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - String: The resulting joined path as an String.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// dir := fs.NewDir("path/to/directory")
|
||||
// joinedPath := dir.Join("subdir", "file.txt")
|
||||
func (d *Dir) Join(elem ...g.String) g.Result[g.String] {
|
||||
path := d.Path()
|
||||
if path.IsErr() {
|
||||
return g.Err[g.String](path.Err())
|
||||
}
|
||||
|
||||
parts := make([]string, len(elem)+1)
|
||||
parts[0] = path.Ok().Std()
|
||||
for i, part := range elem {
|
||||
parts[i+1] = part.Std()
|
||||
}
|
||||
|
||||
return g.Ok(g.String(filepath.Join(parts...)))
|
||||
}
|
||||
|
||||
// SetPath sets the path of the current directory.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - path (String): The new path to be set for the current directory.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - *Dir: A pointer to the updated Dir instance with the new path.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// dir := fs.NewDir("path/to/directory")
|
||||
// dir.SetPath("new/path/to/directory")
|
||||
func (d *Dir) SetPath(path g.String) *Dir {
|
||||
d.path = path
|
||||
return d
|
||||
}
|
||||
|
||||
// CreateAll creates all directories along the given path, with the specified mode (optional).
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - mode ...os.FileMode (optional): The file mode to be used when creating the directories.
|
||||
// If not provided, it defaults to the value of DirDefault constant (0755).
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - *Dir: A pointer to the Dir instance representing the created directories.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// dir := fs.NewDir("path/to/directory")
|
||||
// dir.CreateAll()
|
||||
// dir.CreateAll(0755)
|
||||
func (d *Dir) CreateAll(mode ...os.FileMode) g.Result[*Dir] {
|
||||
path := d.Path()
|
||||
if path.IsErr() {
|
||||
return g.Err[*Dir](path.Err())
|
||||
}
|
||||
|
||||
dmode := os.FileMode(g.DirDefault)
|
||||
if len(mode) > 0 {
|
||||
dmode = mode[0]
|
||||
}
|
||||
|
||||
err := os.MkdirAll(path.Ok().Std(), dmode)
|
||||
if err != nil {
|
||||
return g.Err[*Dir](err)
|
||||
}
|
||||
|
||||
return g.Ok(d)
|
||||
}
|
||||
|
||||
// Rename renames the current directory to the new path.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - newpath String: The new path for the directory.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - Result[*Dir]: A Result containing a pointer to the Dir instance representing
|
||||
// the renamed directory, or an error if the rename fails.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// dir := fs.NewDir("path/to/directory")
|
||||
// dir.Rename("path/to/new_directory")
|
||||
func (d *Dir) Rename(newpath g.String) g.Result[*Dir] {
|
||||
if rd := NewDir(g.String(filepath.Dir(filepath.Clean(newpath.Std())))).CreateAll(); rd.IsErr() {
|
||||
return rd
|
||||
}
|
||||
|
||||
if err := os.Rename(d.path.Std(), newpath.Std()); err != nil {
|
||||
return g.Err[*Dir](err)
|
||||
}
|
||||
|
||||
return g.Ok(NewDir(newpath))
|
||||
}
|
||||
|
||||
// Path returns the absolute path of the current directory.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - Result[String]: The absolute path of the current directory as a String,
|
||||
// or an error if the path cannot be converted to an absolute path.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// dir := fs.NewDir("path/to/directory")
|
||||
// absPath := dir.Path()
|
||||
func (d *Dir) Path() g.Result[g.String] {
|
||||
path, err := filepath.Abs(d.path.Std())
|
||||
if err != nil {
|
||||
return g.Err[g.String](err)
|
||||
}
|
||||
|
||||
return g.Ok(g.String(path))
|
||||
}
|
||||
|
||||
// Exists checks if the current directory exists.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - bool: true if the current directory exists, false otherwise.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// dir := fs.NewDir("path/to/directory")
|
||||
// exists := dir.Exists()
|
||||
func (d *Dir) Exists() bool {
|
||||
path := d.Path()
|
||||
if path.IsErr() {
|
||||
return false
|
||||
}
|
||||
|
||||
info, err := os.Stat(path.Ok().Std())
|
||||
return err == nil && info.IsDir()
|
||||
}
|
||||
|
||||
// Read lazily iterates over the content of the current directory and yields a
|
||||
// File for each entry. Entries are read in bounded batches and retain the
|
||||
// filesystem order; callers that need sorted output can collect and sort them.
|
||||
//
|
||||
// Returns:
|
||||
// - SeqResult[*File]: A sequence of Result[*File] instances representing each file and directory
|
||||
// in the current directory. It returns an error if reading the directory fails.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// dir := fs.NewDir("path/to/directory")
|
||||
// files := dir.Read()
|
||||
// for file := range files {
|
||||
// fmt.Println(file.Ok().Name())
|
||||
// }
|
||||
func (d *Dir) Read() g.SeqResult[*File] {
|
||||
return func(yield func(g.Result[*File]) bool) {
|
||||
dpath := d.Path()
|
||||
if dpath.IsErr() {
|
||||
yield(g.Err[*File](dpath.Err()))
|
||||
return
|
||||
}
|
||||
|
||||
directory, err := os.Open(dpath.Ok().Std())
|
||||
if err != nil {
|
||||
yield(g.Err[*File](err))
|
||||
return
|
||||
}
|
||||
|
||||
defer directory.Close()
|
||||
|
||||
const batchSize = 128
|
||||
base := dpath.Ok().Std()
|
||||
|
||||
for {
|
||||
entries, readErr := directory.ReadDir(batchSize)
|
||||
for _, entry := range entries {
|
||||
if !yield(g.Ok(NewFile(filepath.Join(base, entry.Name())))) {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
switch readErr {
|
||||
case nil:
|
||||
continue
|
||||
case io.EOF:
|
||||
return
|
||||
default:
|
||||
yield(g.Err[*File](readErr))
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Glob iterates over files in the current directory matching a specified pattern and yields File instances for each match.
|
||||
// This method utilizes a lazy evaluation strategy, processing files as they are needed.
|
||||
//
|
||||
// Returns:
|
||||
// - SeqResult[*File]: A sequence of Result[*File] instances representing the files that match the
|
||||
// provided pattern in the current directory. It returns an error if the glob operation fails.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// dir := fs.NewDir("path/to/directory/*.txt")
|
||||
// files := dir.Glob()
|
||||
// for file := range files {
|
||||
// fmt.Println(file.Ok().Name())
|
||||
// }
|
||||
func (d *Dir) Glob() g.SeqResult[*File] {
|
||||
return func(yield func(g.Result[*File]) bool) {
|
||||
matches, err := filepath.Glob(d.path.Std())
|
||||
if err != nil {
|
||||
yield(g.Err[*File](err))
|
||||
return
|
||||
}
|
||||
|
||||
for _, match := range matches {
|
||||
file := NewFile(g.String(match)).Path()
|
||||
if file.IsErr() {
|
||||
yield(g.Err[*File](file.Err()))
|
||||
return
|
||||
}
|
||||
|
||||
if !yield(g.Ok(NewFile(file.Ok()))) {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Walk returns a lazy sequence of all files and directories under the current Dir.
|
||||
// You can customize inclusion/exclusion using SeqResult methods (Exclude, Filter, etc.).
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// NewDir("path/to/dir").
|
||||
// Walk().
|
||||
// Exclude((*File).IsLink).
|
||||
// ForEach(func(r Result[*File]) {
|
||||
// if r.IsOk() {
|
||||
// fmt.Println(r.Ok().Path().Ok().Std())
|
||||
// }
|
||||
// })
|
||||
func (d *Dir) Walk() g.SeqResult[*File] {
|
||||
return func(yield func(g.Result[*File]) bool) {
|
||||
stack := g.SliceOf(d)
|
||||
stopped := false
|
||||
|
||||
// Track resolved directory paths already scheduled for traversal so a
|
||||
// symlink (or hardlinked dir) pointing back into an ancestor does not
|
||||
// drive the stack into unbounded recursion.
|
||||
visited := g.NewSet[g.String]()
|
||||
if root := d.Path(); root.IsOk() {
|
||||
visited.Insert(root.Ok())
|
||||
}
|
||||
|
||||
for !stack.IsEmpty() && !stopped {
|
||||
current := stack.Pop()
|
||||
if current.IsNone() {
|
||||
break
|
||||
}
|
||||
|
||||
current.Some().Read().Range(func(r g.Result[*File]) bool {
|
||||
if r.IsErr() {
|
||||
if !yield(r) {
|
||||
stopped = true
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
file := r.Ok()
|
||||
if !yield(g.Ok(file)) {
|
||||
stopped = true
|
||||
return false
|
||||
}
|
||||
|
||||
// Use Lstat so that a symbolic link to a directory is reported
|
||||
// but not descended into; following it (via Stat) is what makes
|
||||
// symlink cycles loop forever.
|
||||
lstat := file.Lstat()
|
||||
if lstat.IsErr() {
|
||||
if !yield(g.Err[*File](lstat.Err())) {
|
||||
stopped = true
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
if lstat.Ok().Mode()&os.ModeSymlink != 0 {
|
||||
return true
|
||||
}
|
||||
|
||||
if lstat.Ok().IsDir() {
|
||||
path := file.Path()
|
||||
if path.IsErr() {
|
||||
if !yield(g.Err[*File](path.Err())) {
|
||||
stopped = true
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
if visited.Contains(path.Ok()) {
|
||||
return true
|
||||
}
|
||||
|
||||
visited.Insert(path.Ok())
|
||||
stack.Push(NewDir(path.Ok()))
|
||||
}
|
||||
|
||||
return true
|
||||
})
|
||||
|
||||
if stopped {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// String returns the String representation of the current directory's path.
|
||||
func (d *Dir) String() g.String { return d.path }
|
||||
|
||||
// Print writes the content of the Dir to the standard output (console)
|
||||
// and returns the Dir unchanged.
|
||||
func (d *Dir) Print() *Dir { fmt.Print(d); return d }
|
||||
|
||||
// Println writes the content of the Dir to the standard output (console) with a newline
|
||||
// and returns the Dir unchanged.
|
||||
func (d *Dir) Println() *Dir { fmt.Println(d); return d }
|
||||
+155
@@ -0,0 +1,155 @@
|
||||
package fs
|
||||
|
||||
import (
|
||||
"encoding/gob"
|
||||
json "encoding/json/v2"
|
||||
|
||||
"github.com/enetx/g"
|
||||
)
|
||||
|
||||
type (
|
||||
// fencode represents a wrapper for file encoding.
|
||||
fencode struct{ f *File }
|
||||
|
||||
// fdecode represents a wrapper for file decoding.
|
||||
fdecode struct{ f *File }
|
||||
)
|
||||
|
||||
// Encode returns an fencode struct wrapping the given file for encoding.
|
||||
func (f *File) Encode() fencode { return fencode{f} }
|
||||
|
||||
// Decode returns an fdecode struct wrapping the given file for decoding.
|
||||
func (f *File) Decode() fdecode { return fdecode{f} }
|
||||
|
||||
// Gob encodes the provided data using the encoding/gob package and writes it to the file.
|
||||
// It returns a Result[*File] indicating the success or failure of the encoding operation.
|
||||
//
|
||||
// The created file is closed automatically before the method returns.
|
||||
//
|
||||
// Usage:
|
||||
//
|
||||
// data := g.SliceOf(1, 2, 3, 4)
|
||||
// result := fs.NewFile("somefile.gob").Encode().Gob(data)
|
||||
//
|
||||
// Parameters:
|
||||
// - data: The data to be encoded and written to the file.
|
||||
//
|
||||
// Returns:
|
||||
// - Result[*File]: A Result containing a *File if the operation is successful; otherwise, an error Result.
|
||||
func (fe fencode) Gob(data any) g.Result[*File] {
|
||||
r := fe.f.Create()
|
||||
if r.IsErr() {
|
||||
return r
|
||||
}
|
||||
|
||||
defer r.Ok().Close()
|
||||
|
||||
if err := gob.NewEncoder(r.Ok().Std()).Encode(data); err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
// Gob decodes data from the file using the encoding/gob package and populates the provided data structure.
|
||||
// It returns a Result[*File] indicating the success or failure of the decoding operation.
|
||||
//
|
||||
// The file is closed automatically before the method returns.
|
||||
//
|
||||
// Usage:
|
||||
//
|
||||
// var data g.Slice[int]
|
||||
// result := fs.NewFile("somefile.gob").Decode().Gob(&data)
|
||||
//
|
||||
// Parameters:
|
||||
// - data: A pointer to the data structure where the decoded data will be stored.
|
||||
//
|
||||
// Returns:
|
||||
// - Result[*File]: A Result containing a *File if the operation is successful; otherwise, an error Result.
|
||||
func (fd fdecode) Gob(data any) g.Result[*File] {
|
||||
r := fd.f.Open()
|
||||
if r.IsErr() {
|
||||
return r
|
||||
}
|
||||
|
||||
defer r.Ok().Close()
|
||||
|
||||
if err := gob.NewDecoder(r.Ok().Std()).Decode(data); err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
// JSON encodes the provided data using the encoding/json/v2 package and writes it to the file.
|
||||
// It returns a Result[*File] indicating the success or failure of the encoding operation.
|
||||
//
|
||||
// The created file is closed automatically before the method returns.
|
||||
//
|
||||
// Breaking changes (v2 semantics) compared to the previous encoding/json v1 implementation:
|
||||
// - nil slices are marshaled as [] and nil maps as {} (v1 emitted null for both);
|
||||
// - strings containing invalid UTF-8 yield Err (v1 replaced invalid sequences with U+FFFD);
|
||||
// - no trailing newline is written after the JSON value (v1's Encoder.Encode appended '\n');
|
||||
// - '<', '>', '&' and U+2028/U+2029 are emitted raw (v1's Encoder HTML-escaped them).
|
||||
//
|
||||
// Usage:
|
||||
//
|
||||
// data := g.SliceOf(1, 2, 3, 4)
|
||||
// result := fs.NewFile("somefile.json").Encode().JSON(data)
|
||||
//
|
||||
// Parameters:
|
||||
// - data: The data to be encoded and written to the file.
|
||||
//
|
||||
// Returns:
|
||||
// - Result[*File]: A Result containing a *File if the operation is successful; otherwise, an error Result.
|
||||
func (fe fencode) JSON(data any) g.Result[*File] {
|
||||
r := fe.f.Create()
|
||||
if r.IsErr() {
|
||||
return r
|
||||
}
|
||||
|
||||
defer r.Ok().Close()
|
||||
|
||||
if err := json.MarshalWrite(r.Ok().Std(), data); err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
|
||||
// JSON decodes data from the file using the encoding/json/v2 package and populates the provided data structure.
|
||||
// It returns a Result[*File] indicating the success or failure of the decoding operation.
|
||||
//
|
||||
// The file is closed automatically before the method returns.
|
||||
//
|
||||
// Breaking changes (v2 semantics) compared to the previous encoding/json v1 implementation:
|
||||
// - duplicate object member names yield Err (v1 silently kept the last value);
|
||||
// - struct field name matching is case-sensitive (v1 fell back to case-insensitive matching);
|
||||
// - JSON strings containing invalid UTF-8 yield Err (v1 decoded them with U+FFFD replacements);
|
||||
// - the file must contain exactly one JSON value: non-whitespace data after the
|
||||
// top-level value yields Err (v1's Decoder.Decode read one value and ignored the rest).
|
||||
//
|
||||
// Usage:
|
||||
//
|
||||
// var data g.Slice[int]
|
||||
// result := fs.NewFile("somefile.json").Decode().JSON(&data)
|
||||
//
|
||||
// Parameters:
|
||||
// - data: A pointer to the data structure where the decoded data will be stored.
|
||||
//
|
||||
// Returns:
|
||||
// - Result[*File]: A Result containing a *File if the operation is successful; otherwise, an error Result.
|
||||
func (fd fdecode) JSON(data any) g.Result[*File] {
|
||||
r := fd.f.Open()
|
||||
if r.IsErr() {
|
||||
return r
|
||||
}
|
||||
|
||||
defer r.Ok().Close()
|
||||
|
||||
if err := json.UnmarshalRead(r.Ok().Std(), data); err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return r
|
||||
}
|
||||
+17
@@ -0,0 +1,17 @@
|
||||
package fs
|
||||
|
||||
import "fmt"
|
||||
|
||||
// ErrFileNotExist represents an error for when a file does not exist.
|
||||
type ErrFileNotExist struct{ Msg string }
|
||||
|
||||
// Error returns the error message for ErrFileNotExist.
|
||||
func (e *ErrFileNotExist) Error() string { return fmt.Sprintf("no such file: %s", e.Msg) }
|
||||
|
||||
// ErrFileClosed represents an error for when a file is already closed.
|
||||
type ErrFileClosed struct{ Msg string }
|
||||
|
||||
// Error returns the error message for ErrFileClosed.
|
||||
func (e *ErrFileClosed) Error() string {
|
||||
return fmt.Sprintf("%s: file is already closed and unlocked", e.Msg)
|
||||
}
|
||||
+778
@@ -0,0 +1,778 @@
|
||||
// Package fs provides chainable, Result-based filesystem operations (File, Dir),
|
||||
// including lazy SeqResult iterators over file lines and chunks.
|
||||
package fs
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/g/internal/filelock"
|
||||
"github.com/enetx/g/internal/mimesniff"
|
||||
)
|
||||
|
||||
// errChunkSize is the shared error yielded by File.Chunks and File.ChunksRaw
|
||||
// when the requested chunk size is not positive.
|
||||
var errChunkSize = errors.New("chunk size must be > 0")
|
||||
|
||||
// File is a struct that represents a file.
|
||||
type File struct {
|
||||
file *os.File // Underlying os.File.
|
||||
name g.String // File name.
|
||||
guard bool // Guard indicates whether the file is protected against concurrent access.
|
||||
}
|
||||
|
||||
type fileReader struct {
|
||||
owner *File
|
||||
file *os.File
|
||||
}
|
||||
|
||||
func (r fileReader) Read(p []byte) (int, error) { return r.file.Read(p) }
|
||||
func (r fileReader) Close() error {
|
||||
if r.owner.file == r.file {
|
||||
return r.owner.Close()
|
||||
}
|
||||
|
||||
return r.file.Close()
|
||||
}
|
||||
|
||||
// NewFile returns a new File instance with the given name.
|
||||
func NewFile[T ~string](name T) *File { return &File{name: g.String(name)} }
|
||||
|
||||
// Lines returns a new iterator instance that can be used to read the file
|
||||
// line by line.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// // Open a new file with the specified name "text.txt"
|
||||
// fs.NewFile("text.txt").
|
||||
// Lines(). // Read the file line by line
|
||||
// Skip(3). // Skip the first 3 lines
|
||||
// Exclude(f.IsZero). // Exclude empty lines
|
||||
// Dedup(). // Remove consecutive duplicate lines
|
||||
// Map(g.String.Upper). // Convert each line to uppercase
|
||||
// ForEach(func(s g.Result[g.String]) { s.Ok().Print() }) // For each line, print it
|
||||
//
|
||||
// // Output:
|
||||
// // UPPERCASED_LINE4
|
||||
// // UPPERCASED_LINE5
|
||||
// // UPPERCASED_LINE6
|
||||
func (f *File) Lines() g.SeqResult[g.String] {
|
||||
return func(yield func(g.Result[g.String]) bool) {
|
||||
if f.file == nil {
|
||||
if r := f.Open(); r.IsErr() {
|
||||
yield(g.Err[g.String](r.Err()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
defer f.Close()
|
||||
|
||||
scanner := bufio.NewScanner(f.file)
|
||||
scanner.Split(bufio.ScanLines)
|
||||
|
||||
for scanner.Scan() {
|
||||
if !yield(g.Ok(g.String(scanner.Text()))) {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if err := scanner.Err(); err != nil {
|
||||
yield(g.Err[g.String](err))
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// LinesRaw returns a new iterator instance that reads the file line by line,
|
||||
// yielding each line as a Bytes slice (raw []byte).
|
||||
//
|
||||
// This version avoids intermediate string allocations by working directly with byte slices.
|
||||
// The returned Bytes are copies of the scanner buffer and are safe to retain.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - SeqResult[Bytes]: An iterator over raw byte lines from the file.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// fs.NewFile("text.txt").
|
||||
// LinesRaw(). // Read raw byte lines
|
||||
// Filter(func(b g.Bytes) bool {
|
||||
// return len(b) > 0
|
||||
// }).
|
||||
// ForEach(func(line g.Result[g.Bytes]) {
|
||||
// line.Ok().Print()
|
||||
// })
|
||||
//
|
||||
// Output:
|
||||
// LINE_1
|
||||
// LINE_2
|
||||
// ...
|
||||
//
|
||||
// Note: Each line is copied before yielding to avoid scanner buffer reuse issues.
|
||||
func (f *File) LinesRaw() g.SeqResult[g.Bytes] {
|
||||
return func(yield func(g.Result[g.Bytes]) bool) {
|
||||
if f.file == nil {
|
||||
if r := f.Open(); r.IsErr() {
|
||||
yield(g.Err[g.Bytes](r.Err()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
defer f.Close()
|
||||
|
||||
scanner := bufio.NewScanner(f.file)
|
||||
scanner.Split(bufio.ScanLines)
|
||||
|
||||
for scanner.Scan() {
|
||||
line := make(g.Bytes, len(scanner.Bytes()))
|
||||
copy(line, scanner.Bytes())
|
||||
if !yield(g.Ok(line)) {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if err := scanner.Err(); err != nil {
|
||||
yield(g.Err[g.Bytes](err))
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Chunks returns a new iterator instance that can be used to read the file
|
||||
// in fixed-size chunks of the specified size in bytes.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - size (int): The size of each chunk in bytes.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// // Open a new file with the specified name "text.txt"
|
||||
// fs.NewFile("text.txt").
|
||||
// Chunks(100). // Read the file in chunks of 100 bytes
|
||||
// Map(g.String.Upper). // Convert each chunk to uppercase
|
||||
// ForEach(func(s g.Result[g.String]) { s.Ok().Print() }) // For each chunk, print it
|
||||
//
|
||||
// // Output:
|
||||
// // UPPERCASED_CHUNK1
|
||||
// // UPPERCASED_CHUNK2
|
||||
// // UPPERCASED_CHUNK3
|
||||
func (f *File) Chunks(size g.Int) g.SeqResult[g.String] {
|
||||
return func(yield func(g.Result[g.String]) bool) {
|
||||
if size.Lte(0) {
|
||||
yield(g.Err[g.String](errChunkSize))
|
||||
return
|
||||
}
|
||||
|
||||
if f.file == nil {
|
||||
if r := f.Open(); r.IsErr() {
|
||||
yield(g.Err[g.String](r.Err()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
defer f.Close()
|
||||
|
||||
buffer := make([]byte, size)
|
||||
|
||||
for {
|
||||
n, err := f.file.Read(buffer)
|
||||
if err != nil && err != io.EOF {
|
||||
yield(g.Err[g.String](err))
|
||||
return
|
||||
}
|
||||
|
||||
if n == 0 {
|
||||
break
|
||||
}
|
||||
|
||||
if !yield(g.Ok(g.String(buffer[:n]))) {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ChunksRaw returns a new iterator instance that reads the file in fixed-size
|
||||
// chunks of bytes, yielding each chunk as a Bytes slice.
|
||||
//
|
||||
// This method avoids intermediate string allocations and operates directly on byte slices.
|
||||
// Each chunk is copied from the underlying buffer to make it safe for downstream use.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - size (Int): The size of each chunk in bytes. Must be > 0.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - SeqResult[Bytes]: An iterator over raw byte chunks from the file.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// fs.NewFile("text.txt").
|
||||
// ChunksRaw(128). // Read raw 128-byte chunks
|
||||
// ForEach(func(chunk g.Result[g.Bytes]) {
|
||||
// chunk.Ok().Print()
|
||||
// })
|
||||
//
|
||||
// Output:
|
||||
// RAW_CHUNK_1
|
||||
// RAW_CHUNK_2
|
||||
// ...
|
||||
//
|
||||
// Note: Each chunk is copied from the buffer to ensure memory safety.
|
||||
func (f *File) ChunksRaw(size g.Int) g.SeqResult[g.Bytes] {
|
||||
return func(yield func(g.Result[g.Bytes]) bool) {
|
||||
if size.Lte(0) {
|
||||
yield(g.Err[g.Bytes](errChunkSize))
|
||||
return
|
||||
}
|
||||
|
||||
if f.file == nil {
|
||||
if r := f.Open(); r.IsErr() {
|
||||
yield(g.Err[g.Bytes](r.Err()))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
defer f.Close()
|
||||
|
||||
buf := make([]byte, size)
|
||||
|
||||
for {
|
||||
n, err := f.file.Read(buf)
|
||||
if err != nil && err != io.EOF {
|
||||
yield(g.Err[g.Bytes](err))
|
||||
return
|
||||
}
|
||||
|
||||
if n == 0 {
|
||||
break
|
||||
}
|
||||
|
||||
chunk := make(g.Bytes, n)
|
||||
copy(chunk, buf[:n])
|
||||
|
||||
if !yield(g.Ok(chunk)) {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Append appends the given content to the file, with the specified mode (optional).
|
||||
// If no FileMode is provided, the default FileMode (0644) is used.
|
||||
// Don't forget to close the file!
|
||||
func (f *File) Append(content g.String, mode ...os.FileMode) g.Result[*File] {
|
||||
if f.file == nil {
|
||||
if r := f.createAll(); r.IsErr() {
|
||||
return r
|
||||
}
|
||||
|
||||
fmode := os.FileMode(g.FileDefault)
|
||||
if len(mode) > 0 {
|
||||
fmode = mode[0]
|
||||
}
|
||||
|
||||
if r := f.OpenFile(os.O_APPEND|os.O_CREATE|os.O_WRONLY, fmode); r.IsErr() {
|
||||
return r
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := f.file.WriteString(content.Std()); err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return g.Ok(f)
|
||||
}
|
||||
|
||||
// Chmod changes the mode of the file.
|
||||
func (f *File) Chmod(mode os.FileMode) g.Result[*File] {
|
||||
var err error
|
||||
if f.file != nil {
|
||||
err = f.file.Chmod(mode)
|
||||
} else {
|
||||
err = os.Chmod(f.name.Std(), mode)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return g.Ok(f)
|
||||
}
|
||||
|
||||
// Chown changes the owner of the file.
|
||||
func (f *File) Chown(uid, gid int) g.Result[*File] {
|
||||
var err error
|
||||
if f.file != nil {
|
||||
err = f.file.Chown(uid, gid)
|
||||
} else {
|
||||
err = os.Chown(f.name.Std(), uid, gid)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return g.Ok(f)
|
||||
}
|
||||
|
||||
// Seek sets the file offset for the next Read or Write operation. The offset
|
||||
// is specified by the 'offset' parameter, and the 'whence' parameter determines
|
||||
// the reference point for the offset.
|
||||
//
|
||||
// The 'offset' parameter specifies the new offset in bytes relative to the
|
||||
// reference point determined by 'whence'. If 'whence' is set to io.SeekStart,
|
||||
// io.SeekCurrent, or io.SeekEnd, the offset is relative to the start of the file,
|
||||
// the current offset, or the end of the file, respectively.
|
||||
//
|
||||
// If the file is not open, this method will attempt to open it. If the open
|
||||
// operation fails, an error is returned.
|
||||
//
|
||||
// If the Seek operation fails, the file is closed, and an error is returned.
|
||||
//
|
||||
// Example:
|
||||
//
|
||||
// file := fs.NewFile("example.txt")
|
||||
// result := file.Seek(100, io.SeekStart)
|
||||
// if result.Err() != nil {
|
||||
// log.Fatal(result.Err())
|
||||
// }
|
||||
//
|
||||
// Parameters:
|
||||
// - offset: The new offset in bytes.
|
||||
// - whence: The reference point for the offset (io.SeekStart, io.SeekCurrent, or io.SeekEnd).
|
||||
//
|
||||
// Don't forget to close the file!
|
||||
func (f *File) Seek(offset int64, whence int) g.Result[*File] {
|
||||
if f.file == nil {
|
||||
if r := f.Open(); r.IsErr() {
|
||||
return r
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := f.file.Seek(offset, whence); err != nil {
|
||||
f.Close()
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return g.Ok(f)
|
||||
}
|
||||
|
||||
// Close closes the File and unlocks its underlying file, if it is not already closed.
|
||||
func (f *File) Close() error {
|
||||
if f.file == nil {
|
||||
return &ErrFileClosed{f.name.Std()}
|
||||
}
|
||||
|
||||
var err error
|
||||
|
||||
if f.guard {
|
||||
err = filelock.Unlock(f.file)
|
||||
}
|
||||
|
||||
if closeErr := f.file.Close(); closeErr != nil {
|
||||
err = errors.Join(err, closeErr)
|
||||
}
|
||||
|
||||
f.file = nil
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// Copy copies the file to the specified destination, with the specified mode (optional).
|
||||
// If no mode is provided, the default FileMode (0644) is used.
|
||||
func (f *File) Copy(dest g.String, mode ...os.FileMode) g.Result[*File] {
|
||||
if r := f.Open(); r.IsErr() {
|
||||
return r
|
||||
}
|
||||
|
||||
defer f.Close()
|
||||
|
||||
nf := NewFile(dest)
|
||||
if f.guard {
|
||||
nf.guard = true
|
||||
}
|
||||
|
||||
return nf.WriteFromReader(f.file, mode...)
|
||||
}
|
||||
|
||||
// Create is similar to os.Create; if the file is guarded, the returned file is write-locked.
|
||||
// Don't forget to close the file!
|
||||
func (f *File) Create() g.Result[*File] {
|
||||
return f.OpenFile(os.O_RDWR|os.O_CREATE|os.O_TRUNC, g.FileCreate)
|
||||
}
|
||||
|
||||
// Dir returns the directory the file is in as a Dir instance.
|
||||
func (f *File) Dir() g.Result[*Dir] {
|
||||
dirPath := f.dirPath()
|
||||
if dirPath.IsErr() {
|
||||
return g.Err[*Dir](dirPath.Err())
|
||||
}
|
||||
|
||||
return g.Ok(NewDir(dirPath.Ok()))
|
||||
}
|
||||
|
||||
// Exists checks if the file exists.
|
||||
func (f *File) Exists() bool {
|
||||
_, err := os.Stat(f.name.Std())
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// Ext returns the file extension.
|
||||
func (f *File) Ext() g.String { return g.String(filepath.Ext(f.name.Std())) }
|
||||
|
||||
// Guard sets a lock on the file to protect it from concurrent access.
|
||||
// It returns the File instance with the guard enabled.
|
||||
func (f *File) Guard() *File {
|
||||
f.guard = true
|
||||
return f
|
||||
}
|
||||
|
||||
// MimeType returns the MIME type of the file as Result[String].
|
||||
func (f *File) MimeType() g.Result[g.String] {
|
||||
if r := f.Open(); r.IsErr() {
|
||||
return g.Err[g.String](r.Err())
|
||||
}
|
||||
|
||||
defer f.Close()
|
||||
|
||||
buff := make([]byte, mimesniff.SniffLen)
|
||||
|
||||
bytesRead, err := f.file.ReadAt(buff, 0)
|
||||
if err != nil && err != io.EOF {
|
||||
return g.Err[g.String](err)
|
||||
}
|
||||
|
||||
buff = buff[:bytesRead]
|
||||
|
||||
return g.Ok(g.String(mimesniff.DetectContentType(buff)))
|
||||
}
|
||||
|
||||
// Name returns the name of the file.
|
||||
func (f *File) Name() g.String {
|
||||
if f.file != nil {
|
||||
return g.String(filepath.Base(f.file.Name()))
|
||||
}
|
||||
|
||||
return g.String(filepath.Base(f.name.Std()))
|
||||
}
|
||||
|
||||
// Open is like os.Open; if the file is guarded, the returned file is read-locked.
|
||||
// Don't forget to close the file!
|
||||
func (f *File) Open() g.Result[*File] { return f.OpenFile(os.O_RDONLY, 0) }
|
||||
|
||||
// OpenFile is like os.OpenFile; if the file is guarded, the returned file is locked.
|
||||
// If flag includes os.O_WRONLY or os.O_RDWR, the file is write-locked;
|
||||
// otherwise, it is read-locked.
|
||||
// Don't forget to close the file!
|
||||
func (f *File) OpenFile(flag int, perm fs.FileMode) g.Result[*File] {
|
||||
file, err := os.OpenFile(f.name.Std(), flag&^os.O_TRUNC, perm)
|
||||
if err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
if f.guard {
|
||||
switch flag & (os.O_RDONLY | os.O_WRONLY | os.O_RDWR) {
|
||||
case os.O_WRONLY, os.O_RDWR:
|
||||
err = filelock.Lock(file)
|
||||
default:
|
||||
err = filelock.RLock(file)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
file.Close()
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
}
|
||||
|
||||
if flag&os.O_TRUNC == os.O_TRUNC {
|
||||
if err := file.Truncate(0); err != nil {
|
||||
if fi, statErr := file.Stat(); statErr != nil || fi.Mode().IsRegular() {
|
||||
if f.guard {
|
||||
filelock.Unlock(file)
|
||||
}
|
||||
|
||||
file.Close()
|
||||
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Release any descriptor already held by this File before reassigning,
|
||||
// otherwise re-opening (e.g. WriteFromReader on an open handle) leaks the
|
||||
// previous fd and orphans its advisory lock.
|
||||
if f.file != nil {
|
||||
if f.guard {
|
||||
filelock.Unlock(f.file)
|
||||
}
|
||||
|
||||
f.file.Close()
|
||||
}
|
||||
|
||||
f.file = file
|
||||
|
||||
return g.Ok(f)
|
||||
}
|
||||
|
||||
// Path returns the absolute path of the file.
|
||||
func (f *File) Path() g.Result[g.String] { return f.filePath() }
|
||||
|
||||
// Print writes the content of the File to the standard output (console)
|
||||
// and returns the File unchanged.
|
||||
func (f *File) Print() *File { fmt.Print(f.Read().UnwrapOrDefault()); return f }
|
||||
|
||||
// Println writes the content of the File to the standard output (console) with a newline
|
||||
// and returns the File unchanged.
|
||||
func (f *File) Println() *File { fmt.Println(f.Read().UnwrapOrDefault()); return f }
|
||||
|
||||
// Read opens the named file (read-locked if the file is guarded) and returns its contents.
|
||||
func (f *File) Read() g.Result[g.String] {
|
||||
if r := f.Open(); r.IsErr() {
|
||||
return g.Err[g.String](r.Err())
|
||||
}
|
||||
|
||||
defer f.Close()
|
||||
|
||||
content, err := io.ReadAll(f.file)
|
||||
if err != nil {
|
||||
return g.Err[g.String](err)
|
||||
}
|
||||
|
||||
return g.Ok(g.String(content))
|
||||
}
|
||||
|
||||
// Reader returns an io.ReadCloser for reading the file's contents.
|
||||
// If the file is not already open, it attempts to open it automatically.
|
||||
// The caller is responsible for closing the returned reader to release system resources.
|
||||
func (f *File) Reader() g.Result[io.ReadCloser] {
|
||||
if f.file == nil {
|
||||
if r := f.Open(); r.IsErr() {
|
||||
return g.Err[io.ReadCloser](r.Err())
|
||||
}
|
||||
}
|
||||
|
||||
return g.Ok[io.ReadCloser](fileReader{owner: f, file: f.file})
|
||||
}
|
||||
|
||||
// Remove removes the file.
|
||||
func (f *File) Remove() g.Result[*File] {
|
||||
if err := os.Remove(f.name.Std()); err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return g.Ok(f)
|
||||
}
|
||||
|
||||
// Rename renames the file to the specified new path.
|
||||
func (f *File) Rename(newpath g.String) g.Result[*File] {
|
||||
if !f.Exists() {
|
||||
return g.Err[*File](&ErrFileNotExist{f.name.Std()})
|
||||
}
|
||||
|
||||
nf := NewFile(newpath)
|
||||
if f.guard {
|
||||
nf.guard = true
|
||||
}
|
||||
|
||||
if r := nf.createAll(); r.IsErr() {
|
||||
return r
|
||||
}
|
||||
|
||||
if err := os.Rename(f.name.Std(), newpath.Std()); err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return g.Ok(nf)
|
||||
}
|
||||
|
||||
// Split splits the file path into its directory and file components.
|
||||
func (f *File) Split() (*Dir, *File) {
|
||||
path := f.Path()
|
||||
if path.IsErr() {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
dir, file := filepath.Split(path.Ok().Std())
|
||||
|
||||
return NewDir(g.String(dir)), NewFile(g.String(file))
|
||||
}
|
||||
|
||||
// Stat returns the fs.FileInfo of the file.
|
||||
// It calls the file's Stat method if the file is open, or os.Stat otherwise.
|
||||
func (f *File) Stat() g.Result[fs.FileInfo] {
|
||||
if f.file != nil {
|
||||
return g.ResultOf(f.file.Stat())
|
||||
}
|
||||
|
||||
return g.ResultOf(os.Stat(f.name.Std()))
|
||||
}
|
||||
|
||||
// Lstat retrieves information about the symbolic link represented by the *File instance.
|
||||
// It returns a Result[fs.FileInfo] containing details about the symbolic link's metadata.
|
||||
// Unlike Stat, Lstat does not follow the link and provides information about the link itself.
|
||||
func (f *File) Lstat() g.Result[fs.FileInfo] {
|
||||
return g.ResultOf(os.Lstat(f.name.Std()))
|
||||
}
|
||||
|
||||
// IsDir checks if the file is a directory.
|
||||
func (f *File) IsDir() bool {
|
||||
stat := f.Stat()
|
||||
return stat.IsOk() && stat.Ok().IsDir()
|
||||
}
|
||||
|
||||
// IsLink checks if the file is a symbolic link.
|
||||
func (f *File) IsLink() bool {
|
||||
stat := f.Lstat()
|
||||
return stat.IsOk() && stat.Ok().Mode()&os.ModeSymlink != 0
|
||||
}
|
||||
|
||||
// Std returns the underlying *os.File instance.
|
||||
// Don't forget to close the file with Close!
|
||||
func (f *File) Std() *os.File { return f.file }
|
||||
|
||||
// CreateTemp creates a new temporary file in the specified directory with the
|
||||
// specified name pattern and returns a Result, which contains a pointer to the File
|
||||
// or an error if the operation fails.
|
||||
// If no directory is specified, the default directory for temporary files is used.
|
||||
// If no name pattern is specified, the default pattern "*" is used.
|
||||
//
|
||||
// Parameters:
|
||||
//
|
||||
// - args ...String: A variadic parameter specifying the directory and/or name
|
||||
// pattern for the temporary file.
|
||||
//
|
||||
// Returns:
|
||||
//
|
||||
// - *File: A pointer to the File representing the temporary file.
|
||||
//
|
||||
// Example usage:
|
||||
//
|
||||
// tmpfile := fs.CreateTempFile() // Creates a temporary file with default settings
|
||||
// tmpfileWithDir := fs.CreateTempFile("mydir") // Creates a temporary file in "mydir" directory
|
||||
// tmpfileWithPattern := fs.CreateTempFile("", "tmp") // Creates a temporary file with "tmp" pattern
|
||||
//
|
||||
// Call .Guard() on the result to lock it; no guard state is set implicitly.
|
||||
func CreateTempFile(args ...g.String) g.Result[*File] {
|
||||
dir := ""
|
||||
pattern := "*"
|
||||
|
||||
if len(args) != 0 {
|
||||
if len(args) > 1 {
|
||||
pattern = args[1].Std()
|
||||
}
|
||||
|
||||
dir = args[0].Std()
|
||||
}
|
||||
|
||||
tmpfile, err := os.CreateTemp(dir, pattern)
|
||||
if err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
ntmpfile := NewFile(g.String(tmpfile.Name()))
|
||||
ntmpfile.file = tmpfile
|
||||
|
||||
defer ntmpfile.Close()
|
||||
|
||||
return g.Ok(ntmpfile)
|
||||
}
|
||||
|
||||
// Write opens the named file (creating it with the given permissions if needed)
|
||||
// and overwrites it with the given content; if the file is guarded, it is
|
||||
// write-locked while writing.
|
||||
func (f *File) Write(content g.String, mode ...os.FileMode) g.Result[*File] {
|
||||
return f.WriteFromReader(content.Reader(), mode...)
|
||||
}
|
||||
|
||||
// WriteFromReader takes an io.Reader (scr) as input and writes the data from the reader into the file.
|
||||
// If no FileMode is provided, the default FileMode (0644) is used.
|
||||
func (f *File) WriteFromReader(scr io.Reader, mode ...os.FileMode) g.Result[*File] {
|
||||
if f.file == nil {
|
||||
if r := f.createAll(); r.IsErr() {
|
||||
return r
|
||||
}
|
||||
}
|
||||
|
||||
filePath := f.filePath()
|
||||
if filePath.IsErr() {
|
||||
return g.Err[*File](filePath.Err())
|
||||
}
|
||||
|
||||
fmode := os.FileMode(g.FileDefault)
|
||||
if len(mode) > 0 {
|
||||
fmode = mode[0]
|
||||
}
|
||||
|
||||
if r := f.OpenFile(os.O_WRONLY|os.O_CREATE|os.O_TRUNC, fmode); r.IsErr() {
|
||||
return g.Err[*File](r.Err())
|
||||
}
|
||||
|
||||
defer f.Close()
|
||||
|
||||
_, err := io.Copy(f.file, scr)
|
||||
if err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
err = f.file.Sync()
|
||||
if err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
|
||||
return g.Ok(f)
|
||||
}
|
||||
|
||||
// dirPath returns the absolute path of the directory containing the file.
|
||||
func (f *File) dirPath() g.Result[g.String] {
|
||||
name := f.name.Std()
|
||||
if info, err := os.Stat(name); err == nil && info.IsDir() {
|
||||
path, err := filepath.Abs(name)
|
||||
if err != nil {
|
||||
return g.Err[g.String](err)
|
||||
}
|
||||
|
||||
return g.Ok(g.String(path))
|
||||
}
|
||||
|
||||
path, err := filepath.Abs(filepath.Dir(name))
|
||||
if err != nil {
|
||||
return g.Err[g.String](err)
|
||||
}
|
||||
|
||||
return g.Ok(g.String(path))
|
||||
}
|
||||
|
||||
// filePath returns the full file path, including the directory and file name.
|
||||
func (f *File) filePath() g.Result[g.String] {
|
||||
path, err := filepath.Abs(f.name.Std())
|
||||
if err != nil {
|
||||
return g.Err[g.String](err)
|
||||
}
|
||||
|
||||
return g.Ok(g.String(path))
|
||||
}
|
||||
|
||||
func (f *File) createAll() g.Result[*File] {
|
||||
dirPath := f.dirPath()
|
||||
if dirPath.IsErr() {
|
||||
return g.Err[*File](dirPath.Err())
|
||||
}
|
||||
|
||||
if !f.Exists() {
|
||||
if err := os.MkdirAll(dirPath.Ok().Std(), g.DirDefault); err != nil {
|
||||
return g.Err[*File](err)
|
||||
}
|
||||
}
|
||||
|
||||
return g.Ok(f)
|
||||
}
|
||||
+196
@@ -0,0 +1,196 @@
|
||||
package g
|
||||
|
||||
import "math"
|
||||
|
||||
// CheckedAdd adds two Ints, returning None if the addition overflows.
|
||||
func (i Int) CheckedAdd(b Int) Option[Int] {
|
||||
sum := i + b
|
||||
if (b > 0 && sum < i) || (b < 0 && sum > i) {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
return Some(sum)
|
||||
}
|
||||
|
||||
// CheckedSub subtracts b from the Int, returning None if the subtraction overflows.
|
||||
func (i Int) CheckedSub(b Int) Option[Int] {
|
||||
diff := i - b
|
||||
if (b < 0 && diff < i) || (b > 0 && diff > i) {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
return Some(diff)
|
||||
}
|
||||
|
||||
// CheckedMul multiplies two Ints, returning None if the multiplication overflows.
|
||||
func (i Int) CheckedMul(b Int) Option[Int] {
|
||||
if i == 0 || b == 0 {
|
||||
return Some(Int(0))
|
||||
}
|
||||
|
||||
if i == -1 {
|
||||
return b.CheckedNeg()
|
||||
}
|
||||
|
||||
if b == -1 {
|
||||
return i.CheckedNeg()
|
||||
}
|
||||
|
||||
c := i * b
|
||||
if c/b != i {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
return Some(c)
|
||||
}
|
||||
|
||||
// CheckedDiv divides the Int by b, returning None if b is zero or the division overflows.
|
||||
func (i Int) CheckedDiv(b Int) Option[Int] {
|
||||
if b == 0 || (i == math.MinInt && b == -1) {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
return Some(i / b)
|
||||
}
|
||||
|
||||
// CheckedRem computes the remainder of the Int divided by b, returning None if b is zero
|
||||
// or the operation overflows (i == MinInt and b == -1).
|
||||
func (i Int) CheckedRem(b Int) Option[Int] {
|
||||
if b == 0 || (i == math.MinInt && b == -1) {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
return Some(i % b)
|
||||
}
|
||||
|
||||
// CheckedNeg negates the Int, returning None if the negation overflows (i == MinInt).
|
||||
func (i Int) CheckedNeg() Option[Int] {
|
||||
if i == math.MinInt {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
return Some(-i)
|
||||
}
|
||||
|
||||
// CheckedAbs returns the absolute value of the Int, returning None if it overflows (i == MinInt).
|
||||
func (i Int) CheckedAbs() Option[Int] {
|
||||
if i == math.MinInt {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
if i < 0 {
|
||||
return Some(-i)
|
||||
}
|
||||
|
||||
return Some(i)
|
||||
}
|
||||
|
||||
// CheckedPow raises the Int to the power of exp using exponentiation by squaring,
|
||||
// returning None if exp is negative or the computation overflows. An exp of zero yields Some(1).
|
||||
func (i Int) CheckedPow(exp Int) Option[Int] {
|
||||
if exp < 0 {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
result, base := Int(1), i
|
||||
|
||||
for exp > 0 {
|
||||
if exp&1 == 1 {
|
||||
r := result.CheckedMul(base)
|
||||
if r.IsNone() {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
result = r.Some()
|
||||
}
|
||||
|
||||
exp >>= 1
|
||||
|
||||
if exp > 0 {
|
||||
b := base.CheckedMul(base)
|
||||
if b.IsNone() {
|
||||
return None[Int]()
|
||||
}
|
||||
|
||||
base = b.Some()
|
||||
}
|
||||
}
|
||||
|
||||
return Some(result)
|
||||
}
|
||||
|
||||
// SaturatingAdd adds two Ints, clamping the result to MinInt or MaxInt on overflow.
|
||||
func (i Int) SaturatingAdd(b Int) Int {
|
||||
sum := i + b
|
||||
if b > 0 && sum < i {
|
||||
return math.MaxInt
|
||||
}
|
||||
|
||||
if b < 0 && sum > i {
|
||||
return math.MinInt
|
||||
}
|
||||
|
||||
return sum
|
||||
}
|
||||
|
||||
// SaturatingSub subtracts b from the Int, clamping the result to MinInt or MaxInt on overflow.
|
||||
func (i Int) SaturatingSub(b Int) Int {
|
||||
diff := i - b
|
||||
if b < 0 && diff < i {
|
||||
return math.MaxInt
|
||||
}
|
||||
|
||||
if b > 0 && diff > i {
|
||||
return math.MinInt
|
||||
}
|
||||
|
||||
return diff
|
||||
}
|
||||
|
||||
// SaturatingMul multiplies two Ints, clamping the result to MinInt or MaxInt on overflow.
|
||||
func (i Int) SaturatingMul(b Int) Int {
|
||||
if c := i.CheckedMul(b); c.IsSome() {
|
||||
return c.Some()
|
||||
}
|
||||
|
||||
if (i > 0) == (b > 0) {
|
||||
return math.MaxInt
|
||||
}
|
||||
|
||||
return math.MinInt
|
||||
}
|
||||
|
||||
// OverflowingAdd adds two Ints, returning the wrapped result and a flag indicating overflow.
|
||||
func (i Int) OverflowingAdd(b Int) (Int, bool) {
|
||||
sum := i + b
|
||||
|
||||
return sum, (b > 0 && sum < i) || (b < 0 && sum > i)
|
||||
}
|
||||
|
||||
// OverflowingSub subtracts b from the Int, returning the wrapped result and a flag indicating overflow.
|
||||
func (i Int) OverflowingSub(b Int) (Int, bool) {
|
||||
diff := i - b
|
||||
|
||||
return diff, (b < 0 && diff < i) || (b > 0 && diff > i)
|
||||
}
|
||||
|
||||
// OverflowingMul multiplies two Ints, returning the wrapped result and a flag indicating overflow.
|
||||
func (i Int) OverflowingMul(b Int) (Int, bool) {
|
||||
return i * b, i.CheckedMul(b).IsNone()
|
||||
}
|
||||
|
||||
// Clamp restricts the Int to the inclusive range [min, max].
|
||||
// The caller must ensure min <= max: this method does not
|
||||
// panic on an inverted range — the lower bound is checked first, so the result
|
||||
// is unspecified when min > max.
|
||||
func (i Int) Clamp(min, max Int) Int {
|
||||
if i < min {
|
||||
return min
|
||||
}
|
||||
|
||||
if i > max {
|
||||
return max
|
||||
}
|
||||
|
||||
return i
|
||||
}
|
||||
+314
@@ -0,0 +1,314 @@
|
||||
// Copyright 2011 The Go Authors. All rights reserved.
|
||||
|
||||
// Package mimesniff is a verbatim copy of the Go standard library's content
|
||||
// sniffer (`net/http/internal/sniff.go`), reproduced here under the Go BSD-3
|
||||
// licence; see the Go Authors copyright header at the top of this file.
|
||||
//
|
||||
// It exists so that fs.File.MimeType does not have to import net/http. That
|
||||
// one import would pull crypto/tls, crypto/x509 and the whole FIPS-140 tree
|
||||
// into the dependency closure of the fs package and everything that imports
|
||||
// it. The sniffer itself needs only bytes and encoding/binary.
|
||||
//
|
||||
// Keep this file a verbatim copy: it implements https://mimesniff.spec.whatwg.org/
|
||||
// and any local edit silently diverges File.MimeType from net/http's result.
|
||||
// To refresh it, re-copy from GOROOT and re-apply only the package clause.
|
||||
package mimesniff
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
)
|
||||
|
||||
// The algorithm uses at most SniffLen bytes to make its decision.
|
||||
const SniffLen = 512
|
||||
|
||||
// DetectContentType implements the algorithm described
|
||||
// at https://mimesniff.spec.whatwg.org/ to determine the
|
||||
// Content-Type of the given data. It considers at most the
|
||||
// first 512 bytes of data. DetectContentType always returns
|
||||
// a valid MIME type: if it cannot determine a more specific one, it
|
||||
// returns "application/octet-stream".
|
||||
func DetectContentType(data []byte) string {
|
||||
if len(data) > SniffLen {
|
||||
data = data[:SniffLen]
|
||||
}
|
||||
|
||||
// Index of the first non-whitespace byte in data.
|
||||
firstNonWS := 0
|
||||
for ; firstNonWS < len(data) && isWS(data[firstNonWS]); firstNonWS++ {
|
||||
}
|
||||
|
||||
for _, sig := range sniffSignatures {
|
||||
if ct := sig.match(data, firstNonWS); ct != "" {
|
||||
return ct
|
||||
}
|
||||
}
|
||||
|
||||
return "application/octet-stream" // fallback
|
||||
}
|
||||
|
||||
// isWS reports whether the provided byte is a whitespace byte (0xWS)
|
||||
// as defined in https://mimesniff.spec.whatwg.org/#terminology.
|
||||
func isWS(b byte) bool {
|
||||
switch b {
|
||||
case '\t', '\n', '\x0c', '\r', ' ':
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// isTT reports whether the provided byte is a tag-terminating byte (0xTT)
|
||||
// as defined in https://mimesniff.spec.whatwg.org/#terminology.
|
||||
func isTT(b byte) bool {
|
||||
switch b {
|
||||
case ' ', '>':
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
type sniffSig interface {
|
||||
// match returns the MIME type of the data, or "" if unknown.
|
||||
match(data []byte, firstNonWS int) string
|
||||
}
|
||||
|
||||
// Data matching the table in section 6.
|
||||
var sniffSignatures = []sniffSig{
|
||||
htmlSig("<!DOCTYPE HTML"),
|
||||
htmlSig("<HTML"),
|
||||
htmlSig("<HEAD"),
|
||||
htmlSig("<SCRIPT"),
|
||||
htmlSig("<IFRAME"),
|
||||
htmlSig("<H1"),
|
||||
htmlSig("<DIV"),
|
||||
htmlSig("<FONT"),
|
||||
htmlSig("<TABLE"),
|
||||
htmlSig("<A"),
|
||||
htmlSig("<STYLE"),
|
||||
htmlSig("<TITLE"),
|
||||
htmlSig("<B"),
|
||||
htmlSig("<BODY"),
|
||||
htmlSig("<BR"),
|
||||
htmlSig("<P"),
|
||||
htmlSig("<!--"),
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\xFF\xFF\xFF"),
|
||||
pat: []byte("<?xml"),
|
||||
skipWS: true,
|
||||
ct: "text/xml; charset=utf-8"},
|
||||
&exactSig{[]byte("%PDF-"), "application/pdf"},
|
||||
&exactSig{[]byte("%!PS-Adobe-"), "application/postscript"},
|
||||
|
||||
// UTF BOMs.
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\x00\x00"),
|
||||
pat: []byte("\xFE\xFF\x00\x00"),
|
||||
ct: "text/plain; charset=utf-16be",
|
||||
},
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\x00\x00"),
|
||||
pat: []byte("\xFF\xFE\x00\x00"),
|
||||
ct: "text/plain; charset=utf-16le",
|
||||
},
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\xFF\x00"),
|
||||
pat: []byte("\xEF\xBB\xBF\x00"),
|
||||
ct: "text/plain; charset=utf-8",
|
||||
},
|
||||
|
||||
// Image types
|
||||
// For posterity, we originally returned "image/vnd.microsoft.icon" from
|
||||
// https://tools.ietf.org/html/draft-ietf-websec-mime-sniff-03#section-7
|
||||
// https://codereview.appspot.com/4746042
|
||||
// but that has since been replaced with "image/x-icon" in Section 6.2
|
||||
// of https://mimesniff.spec.whatwg.org/#matching-an-image-type-pattern
|
||||
&exactSig{[]byte("\x00\x00\x01\x00"), "image/x-icon"},
|
||||
&exactSig{[]byte("\x00\x00\x02\x00"), "image/x-icon"},
|
||||
&exactSig{[]byte("BM"), "image/bmp"},
|
||||
&exactSig{[]byte("GIF87a"), "image/gif"},
|
||||
&exactSig{[]byte("GIF89a"), "image/gif"},
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\xFF\xFF\x00\x00\x00\x00\xFF\xFF\xFF\xFF\xFF\xFF"),
|
||||
pat: []byte("RIFF\x00\x00\x00\x00WEBPVP"),
|
||||
ct: "image/webp",
|
||||
},
|
||||
&exactSig{[]byte("\x89PNG\x0D\x0A\x1A\x0A"), "image/png"},
|
||||
&exactSig{[]byte("\xFF\xD8\xFF"), "image/jpeg"},
|
||||
|
||||
// Audio and Video types
|
||||
// Enforce the pattern match ordering as prescribed in
|
||||
// https://mimesniff.spec.whatwg.org/#matching-an-audio-or-video-type-pattern
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\xFF\xFF\x00\x00\x00\x00\xFF\xFF\xFF\xFF"),
|
||||
pat: []byte("FORM\x00\x00\x00\x00AIFF"),
|
||||
ct: "audio/aiff",
|
||||
},
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\xFF"),
|
||||
pat: []byte("ID3"),
|
||||
ct: "audio/mpeg",
|
||||
},
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\xFF\xFF\xFF"),
|
||||
pat: []byte("OggS\x00"),
|
||||
ct: "application/ogg",
|
||||
},
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\xFF\xFF\xFF\xFF\xFF\xFF"),
|
||||
pat: []byte("MThd\x00\x00\x00\x06"),
|
||||
ct: "audio/midi",
|
||||
},
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\xFF\xFF\x00\x00\x00\x00\xFF\xFF\xFF\xFF"),
|
||||
pat: []byte("RIFF\x00\x00\x00\x00AVI "),
|
||||
ct: "video/avi",
|
||||
},
|
||||
&maskedSig{
|
||||
mask: []byte("\xFF\xFF\xFF\xFF\x00\x00\x00\x00\xFF\xFF\xFF\xFF"),
|
||||
pat: []byte("RIFF\x00\x00\x00\x00WAVE"),
|
||||
ct: "audio/wave",
|
||||
},
|
||||
// 6.2.0.2. video/mp4
|
||||
mp4Sig{},
|
||||
// 6.2.0.3. video/webm
|
||||
&exactSig{[]byte("\x1A\x45\xDF\xA3"), "video/webm"},
|
||||
|
||||
// Font types
|
||||
&maskedSig{
|
||||
// 34 NULL bytes followed by the string "LP"
|
||||
pat: []byte("\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00LP"),
|
||||
// 34 NULL bytes followed by \xF\xF
|
||||
mask: []byte("\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\xFF\xFF"),
|
||||
ct: "application/vnd.ms-fontobject",
|
||||
},
|
||||
&exactSig{[]byte("\x00\x01\x00\x00"), "font/ttf"},
|
||||
&exactSig{[]byte("OTTO"), "font/otf"},
|
||||
&exactSig{[]byte("ttcf"), "font/collection"},
|
||||
&exactSig{[]byte("wOFF"), "font/woff"},
|
||||
&exactSig{[]byte("wOF2"), "font/woff2"},
|
||||
|
||||
// Archive types
|
||||
&exactSig{[]byte("\x1F\x8B\x08"), "application/x-gzip"},
|
||||
&exactSig{[]byte("PK\x03\x04"), "application/zip"},
|
||||
// RAR's signatures are incorrectly defined by the MIME spec as per
|
||||
// https://github.com/whatwg/mimesniff/issues/63
|
||||
// However, RAR Labs correctly defines it at:
|
||||
// https://www.rarlab.com/technote.htm#rarsign
|
||||
// so we use the definition from RAR Labs.
|
||||
// TODO: do whatever the spec ends up doing.
|
||||
&exactSig{[]byte("Rar!\x1A\x07\x00"), "application/x-rar-compressed"}, // RAR v1.5-v4.0
|
||||
&exactSig{[]byte("Rar!\x1A\x07\x01\x00"), "application/x-rar-compressed"}, // RAR v5+
|
||||
|
||||
&exactSig{[]byte("\x00\x61\x73\x6D"), "application/wasm"},
|
||||
|
||||
textSig{}, // should be last
|
||||
}
|
||||
|
||||
type exactSig struct {
|
||||
sig []byte
|
||||
ct string
|
||||
}
|
||||
|
||||
func (e *exactSig) match(data []byte, firstNonWS int) string {
|
||||
if bytes.HasPrefix(data, e.sig) {
|
||||
return e.ct
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type maskedSig struct {
|
||||
mask, pat []byte
|
||||
skipWS bool
|
||||
ct string
|
||||
}
|
||||
|
||||
func (m *maskedSig) match(data []byte, firstNonWS int) string {
|
||||
// pattern matching algorithm section 6
|
||||
// https://mimesniff.spec.whatwg.org/#pattern-matching-algorithm
|
||||
|
||||
if m.skipWS {
|
||||
data = data[firstNonWS:]
|
||||
}
|
||||
if len(m.pat) != len(m.mask) {
|
||||
return ""
|
||||
}
|
||||
if len(data) < len(m.pat) {
|
||||
return ""
|
||||
}
|
||||
for i, pb := range m.pat {
|
||||
maskedData := data[i] & m.mask[i]
|
||||
if maskedData != pb {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
return m.ct
|
||||
}
|
||||
|
||||
type htmlSig []byte
|
||||
|
||||
func (h htmlSig) match(data []byte, firstNonWS int) string {
|
||||
data = data[firstNonWS:]
|
||||
if len(data) < len(h)+1 {
|
||||
return ""
|
||||
}
|
||||
for i, b := range h {
|
||||
db := data[i]
|
||||
if 'A' <= b && b <= 'Z' {
|
||||
db &= 0xDF
|
||||
}
|
||||
if b != db {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
// Next byte must be a tag-terminating byte(0xTT).
|
||||
if !isTT(data[len(h)]) {
|
||||
return ""
|
||||
}
|
||||
return "text/html; charset=utf-8"
|
||||
}
|
||||
|
||||
var mp4ftype = []byte("ftyp")
|
||||
var mp4 = []byte("mp4")
|
||||
|
||||
type mp4Sig struct{}
|
||||
|
||||
func (mp4Sig) match(data []byte, firstNonWS int) string {
|
||||
// https://mimesniff.spec.whatwg.org/#signature-for-mp4
|
||||
// c.f. section 6.2.1
|
||||
if len(data) < 12 {
|
||||
return ""
|
||||
}
|
||||
boxSize := int(binary.BigEndian.Uint32(data[:4]))
|
||||
if len(data) < boxSize || boxSize%4 != 0 {
|
||||
return ""
|
||||
}
|
||||
if !bytes.Equal(data[4:8], mp4ftype) {
|
||||
return ""
|
||||
}
|
||||
for st := 8; st < boxSize; st += 4 {
|
||||
if st == 12 {
|
||||
// Ignores the four bytes that correspond to the version number of the "major brand".
|
||||
continue
|
||||
}
|
||||
if bytes.Equal(data[st:st+3], mp4) {
|
||||
return "video/mp4"
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type textSig struct{}
|
||||
|
||||
func (textSig) match(data []byte, firstNonWS int) string {
|
||||
// c.f. section 5, step 4.
|
||||
for _, b := range data[firstNonWS:] {
|
||||
switch {
|
||||
case b <= 0x08,
|
||||
b == 0x0B,
|
||||
0x0E <= b && b <= 0x1A,
|
||||
0x1C <= b && b <= 0x1F:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
return "text/plain; charset=utf-8"
|
||||
}
|
||||
+80
@@ -0,0 +1,80 @@
|
||||
package g
|
||||
|
||||
// Conversions between the three map types live here as free functions rather
|
||||
// than methods. As methods they weld Map, MapOrd and MapSafe into one
|
||||
// instantiation cluster: any package touching one of them would compile the
|
||||
// methods of all three (and everything those methods mention, transitively).
|
||||
// As free functions each conversion costs only the packages that call it.
|
||||
|
||||
// MapOrdFromMap converts a standard Map to an ordered Map.
|
||||
func MapOrdFromMap[K comparable, V any](m Map[K, V]) MapOrd[K, V] {
|
||||
mo := make(MapOrd[K, V], 0, len(m))
|
||||
for k, v := range m {
|
||||
mo = append(mo, Pair[K, V]{Key: k, Value: v})
|
||||
}
|
||||
|
||||
return mo
|
||||
}
|
||||
|
||||
// MapSafeFromMap converts a standard Map to a thread-safe Map.
|
||||
func MapSafeFromMap[K comparable, V any](m Map[K, V]) *MapSafe[K, V] {
|
||||
ms := NewMapSafe[K, V]()
|
||||
for k, v := range m {
|
||||
ms.Insert(k, v)
|
||||
}
|
||||
|
||||
return ms
|
||||
}
|
||||
|
||||
// MapFromMapOrd converts an ordered Map to a standard Map.
|
||||
func MapFromMapOrd[K comparable, V any](mo MapOrd[K, V]) Map[K, V] {
|
||||
m := NewMap[K, V](mo.Len())
|
||||
for _, p := range mo {
|
||||
m[p.Key] = p.Value
|
||||
}
|
||||
|
||||
return m
|
||||
}
|
||||
|
||||
// MapSafeFromMapOrd converts an ordered Map to a thread-safe Map.
|
||||
func MapSafeFromMapOrd[K comparable, V any](mo MapOrd[K, V]) *MapSafe[K, V] {
|
||||
ms := NewMapSafe[K, V]()
|
||||
for _, p := range mo {
|
||||
ms.Insert(p.Unpack())
|
||||
}
|
||||
|
||||
return ms
|
||||
}
|
||||
|
||||
// MapFromMapSafe converts the MapSafe to a standard Map by taking a snapshot of its
|
||||
// current key-value pairs.
|
||||
//
|
||||
// The returned Map is an independent, non-thread-safe copy; subsequent
|
||||
// mutations to the MapSafe are not reflected in it.
|
||||
func MapFromMapSafe[K comparable, V any](ms *MapSafe[K, V]) Map[K, V] {
|
||||
m := NewMap[K, V](ms.Len())
|
||||
|
||||
ms.data.Range(func(key, value any) bool {
|
||||
m[key.(K)] = *(value.(*V))
|
||||
return true
|
||||
})
|
||||
|
||||
return m
|
||||
}
|
||||
|
||||
// MapOrdFromMapSafe converts the MapSafe to an ordered Map by taking a snapshot of its
|
||||
// current key-value pairs.
|
||||
//
|
||||
// Because MapSafe does not track insertion order, the order of the returned
|
||||
// MapOrd is unspecified. The returned MapOrd is an independent, non-thread-safe
|
||||
// copy; subsequent mutations to the MapSafe are not reflected in it.
|
||||
func MapOrdFromMapSafe[K comparable, V any](ms *MapSafe[K, V]) MapOrd[K, V] {
|
||||
mo := NewMapOrd[K, V](ms.Len())
|
||||
|
||||
ms.data.Range(func(key, value any) bool {
|
||||
mo = append(mo, Pair[K, V]{Key: key.(K), Value: *(value.(*V))})
|
||||
return true
|
||||
})
|
||||
|
||||
return mo
|
||||
}
|
||||
+68
@@ -0,0 +1,68 @@
|
||||
package g
|
||||
|
||||
import "sync"
|
||||
|
||||
// Mutex is a mutual exclusion lock that protects a value of type T.
|
||||
// Unlike sync.Mutex, it binds the protected data to the lock itself,
|
||||
// making it impossible to access the data without holding the lock.
|
||||
//
|
||||
// A Mutex must not be copied after first use: it embeds a sync.Mutex and a
|
||||
// copy would protect a different value than the original (go vet's copylocks
|
||||
// analyzer flags such copies). Always pass a *Mutex, never a Mutex by value.
|
||||
type Mutex[T any] struct {
|
||||
mu sync.Mutex
|
||||
val T
|
||||
}
|
||||
|
||||
// MutexGuard provides access to the value protected by a Mutex.
|
||||
// The guard must be explicitly unlocked when done.
|
||||
//
|
||||
// A guard holds pointers into its owning Mutex; copying the guard and calling
|
||||
// Unlock on more than one copy unlocks the same underlying lock twice, which
|
||||
// panics. Use each guard exactly once and do not copy it.
|
||||
type MutexGuard[T any] struct {
|
||||
mu *sync.Mutex
|
||||
val *T
|
||||
}
|
||||
|
||||
// NewMutex creates a new Mutex containing the given value.
|
||||
func NewMutex[T any](value T) *Mutex[T] { return &Mutex[T]{val: value} }
|
||||
|
||||
// Lock acquires the mutex and returns a guard that provides access to the protected value.
|
||||
// The caller must call Unlock on the guard when done (typically via defer).
|
||||
func (m *Mutex[T]) Lock() MutexGuard[T] {
|
||||
m.mu.Lock()
|
||||
return MutexGuard[T]{mu: &m.mu, val: &m.val}
|
||||
}
|
||||
|
||||
// With acquires the mutex, calls fn with a pointer to the protected value,
|
||||
// and releases the mutex when fn returns.
|
||||
// This is a convenience method that eliminates the need for manual Lock/Unlock management.
|
||||
func (m *Mutex[T]) With(fn func(*T)) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
fn(&m.val)
|
||||
}
|
||||
|
||||
// TryLock attempts to acquire the mutex without blocking.
|
||||
// Returns Some(guard) if successful, None if the mutex is already locked.
|
||||
func (m *Mutex[T]) TryLock() Option[MutexGuard[T]] {
|
||||
if m.mu.TryLock() {
|
||||
return Some(MutexGuard[T]{mu: &m.mu, val: &m.val})
|
||||
}
|
||||
|
||||
return None[MutexGuard[T]]()
|
||||
}
|
||||
|
||||
// Get returns a copy of the protected value.
|
||||
func (g MutexGuard[T]) Get() T { return *g.val }
|
||||
|
||||
// Set replaces the protected value with a new one.
|
||||
func (g MutexGuard[T]) Set(value T) { *g.val = value }
|
||||
|
||||
// Deref returns a pointer to the protected value for direct manipulation.
|
||||
func (g MutexGuard[T]) Deref() *T { return g.val }
|
||||
|
||||
// Unlock releases the mutex. Must be called when done with the guard.
|
||||
func (g MutexGuard[T]) Unlock() { g.mu.Unlock() }
|
||||
+91
@@ -0,0 +1,91 @@
|
||||
package g
|
||||
|
||||
import (
|
||||
"encoding/json/jsontext"
|
||||
json "encoding/json/v2"
|
||||
)
|
||||
|
||||
// jsonNull is the JSON literal returned when marshaling a None Option.
|
||||
// It is a package-level value to avoid allocating a new []byte on every None marshal.
|
||||
//
|
||||
// NOTE: callers (encoding/json) must not mutate the returned slice; the standard
|
||||
// library treats Marshaler output as read-only, so sharing this backing array is safe.
|
||||
var jsonNull = []byte("null")
|
||||
|
||||
// MarshalJSON implements the json.Marshaler interface (encoding/json v1) for Option[T].
|
||||
// Some(value) is marshaled as the JSON representation of value.
|
||||
// None is marshaled as null.
|
||||
//
|
||||
// BREAKING: the implementation is backed by encoding/json/v2, which changes some
|
||||
// edge-case semantics compared to the previous encoding/json implementation:
|
||||
// - Some of a nil slice marshals as [] and Some of a nil map marshals as {},
|
||||
// rather than null.
|
||||
// - Strings containing invalid UTF-8 are rejected with an error instead of
|
||||
// being silently replaced with U+FFFD.
|
||||
func (o Option[T]) MarshalJSON() ([]byte, error) {
|
||||
if o.IsNone() {
|
||||
return jsonNull, nil
|
||||
}
|
||||
|
||||
return json.Marshal(o.v)
|
||||
}
|
||||
|
||||
// UnmarshalJSON implements the json.Unmarshaler interface (encoding/json v1) for Option[T].
|
||||
// JSON null is unmarshaled as None.
|
||||
// Any other valid JSON value is unmarshaled as Some(value).
|
||||
//
|
||||
// BREAKING: the implementation is backed by encoding/json/v2, which is stricter
|
||||
// than the previous encoding/json implementation: duplicate object keys inside
|
||||
// the value are rejected, struct field names match case-sensitively, and strings
|
||||
// containing invalid UTF-8 are rejected.
|
||||
func (o *Option[T]) UnmarshalJSON(data []byte) error {
|
||||
if string(data) == "null" {
|
||||
*o = None[T]()
|
||||
return nil
|
||||
}
|
||||
|
||||
var v T
|
||||
if err := json.Unmarshal(data, &v); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
*o = Some(v)
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarshalJSONTo implements the json.MarshalerTo interface (encoding/json/v2) for Option[T].
|
||||
// encoding/json/v2 prefers this method over MarshalJSON.
|
||||
// Some(value) is encoded as the JSON representation of value; None is encoded as null.
|
||||
//
|
||||
// Because None and JSON null share one representation, nested Options collapse:
|
||||
// Some(None) marshals to null and unmarshals back as None. Wrap the inner
|
||||
// Option in a struct (or use Result) when the distinction must survive a round trip.
|
||||
func (o Option[T]) MarshalJSONTo(enc *jsontext.Encoder) error {
|
||||
if o.IsNone() {
|
||||
return enc.WriteToken(jsontext.Null)
|
||||
}
|
||||
|
||||
return json.MarshalEncode(enc, o.v)
|
||||
}
|
||||
|
||||
// UnmarshalJSONFrom implements the json.UnmarshalerFrom interface (encoding/json/v2) for Option[T].
|
||||
// encoding/json/v2 prefers this method over UnmarshalJSON.
|
||||
// JSON null is decoded as None; any other value is decoded into T as Some(value).
|
||||
func (o *Option[T]) UnmarshalJSONFrom(dec *jsontext.Decoder) error {
|
||||
if dec.PeekKind() == 'n' {
|
||||
if _, err := dec.ReadToken(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
*o = None[T]()
|
||||
return nil
|
||||
}
|
||||
|
||||
var v T
|
||||
if err := json.UnmarshalDecode(dec, &v); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
*o = Some(v)
|
||||
return nil
|
||||
}
|
||||
+239
@@ -0,0 +1,239 @@
|
||||
package g
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"database/sql/driver"
|
||||
"fmt"
|
||||
"math"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Scan implements the database/sql.Scanner interface for Option[T].
|
||||
//
|
||||
// Behavior:
|
||||
// - If src is nil, the Option is set to None (SQL NULL).
|
||||
// - If T implements sql.Scanner, its Scan method is used.
|
||||
// - If src can be directly assigned to T, it is assigned as-is.
|
||||
// - Otherwise, common database type conversions are attempted (e.g., int64 → int, []byte → string).
|
||||
//
|
||||
// Supported conversions (common SQL types):
|
||||
// - INTEGER → int, int8, int16, int32, int64, uint*
|
||||
// - REAL → float32, float64
|
||||
// - TEXT → string, []byte
|
||||
// - BLOB → []byte
|
||||
// - BOOLEAN → bool
|
||||
// - TIMESTAMP → time.Time
|
||||
//
|
||||
// Driver-owned []byte buffers (BLOB/TEXT) are copied before being stored, so the
|
||||
// scanned Option keeps a value that is safe to retain across subsequent rows
|
||||
// (per the database/sql Scanner contract).
|
||||
//
|
||||
// Returns an error if the value cannot be converted to T.
|
||||
func (o *Option[T]) Scan(src any) error {
|
||||
if src == nil {
|
||||
*o = None[T]()
|
||||
return nil
|
||||
}
|
||||
|
||||
var v T
|
||||
|
||||
if scanner, ok := any(&v).(sql.Scanner); ok {
|
||||
if err := scanner.Scan(src); err != nil {
|
||||
return err
|
||||
}
|
||||
*o = Some(v)
|
||||
return nil
|
||||
}
|
||||
|
||||
if val, ok := src.(T); ok {
|
||||
// Copy driver-owned []byte buffers before storing: database/sql may
|
||||
// overwrite the backing array on the next row (database/sql contract).
|
||||
if b, isBytes := any(val).([]byte); isBytes {
|
||||
val = any(append([]byte(nil), b...)).(T)
|
||||
}
|
||||
*o = Some(val)
|
||||
return nil
|
||||
}
|
||||
|
||||
if converted, ok := convertToT[T](src); ok {
|
||||
*o = Some(converted)
|
||||
return nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("Option.Scan: cannot scan %T into %T", src, v)
|
||||
}
|
||||
|
||||
// Value implements the database/sql/driver.Valuer interface for Option[T].
|
||||
//
|
||||
// Behavior:
|
||||
// - If the Option is None, returns nil (SQL NULL).
|
||||
// - If T implements driver.Valuer, its Value method is used.
|
||||
// - If the underlying value is already a valid driver.Value type (int64, float64, bool, []byte, string, time.Time), it is returned directly.
|
||||
// - Otherwise, safe conversions are applied (int → int64, uint → int64, float32 → float64).
|
||||
//
|
||||
// Returns an error if the value cannot be converted to a driver.Value.
|
||||
func (o Option[T]) Value() (driver.Value, error) {
|
||||
if o.IsNone() {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if valuer, ok := any(o.v).(driver.Valuer); ok {
|
||||
return valuer.Value()
|
||||
}
|
||||
|
||||
switch val := any(o.v).(type) {
|
||||
case int64, float64, bool, []byte, string, time.Time:
|
||||
return val, nil
|
||||
default:
|
||||
if converted, ok := convertToDriverValue(val); ok {
|
||||
return converted, nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil, fmt.Errorf("Option.Value: unsupported type %T", o.v)
|
||||
}
|
||||
|
||||
// convertToT attempts to safely convert a source value from database/sql
|
||||
// into type T, supporting common database types without using reflection.
|
||||
//
|
||||
// Only standard Go primitive types are supported:
|
||||
// - int64 → int, int8, int16, int32, uint*, uint64 (if fits)
|
||||
// - float64 → float32, float64
|
||||
// - string / []byte → string
|
||||
// - []byte / string → []byte
|
||||
// - bool → bool
|
||||
// - time.Time → time.Time
|
||||
//
|
||||
// Returns the converted value and true on success, otherwise the zero value of T and false.
|
||||
func convertToT[T any](src any) (T, bool) {
|
||||
var zero T
|
||||
|
||||
switch any(zero).(type) {
|
||||
case int:
|
||||
if i64, ok := src.(int64); ok && fitsInt(i64) {
|
||||
return any(int(i64)).(T), true
|
||||
}
|
||||
case int8:
|
||||
if i64, ok := src.(int64); ok && fitsInt8(i64) {
|
||||
return any(int8(i64)).(T), true
|
||||
}
|
||||
case int16:
|
||||
if i64, ok := src.(int64); ok && fitsInt16(i64) {
|
||||
return any(int16(i64)).(T), true
|
||||
}
|
||||
case int32:
|
||||
if i64, ok := src.(int64); ok && fitsInt32(i64) {
|
||||
return any(int32(i64)).(T), true
|
||||
}
|
||||
case int64:
|
||||
if i64, ok := src.(int64); ok {
|
||||
return any(i64).(T), true
|
||||
}
|
||||
case uint:
|
||||
if i64, ok := src.(int64); ok && i64 >= 0 {
|
||||
return any(uint(i64)).(T), true
|
||||
}
|
||||
case uint8:
|
||||
if i64, ok := src.(int64); ok && i64 >= 0 && i64 <= math.MaxUint8 {
|
||||
return any(uint8(i64)).(T), true
|
||||
}
|
||||
case uint16:
|
||||
if i64, ok := src.(int64); ok && i64 >= 0 && i64 <= math.MaxUint16 {
|
||||
return any(uint16(i64)).(T), true
|
||||
}
|
||||
case uint32:
|
||||
if i64, ok := src.(int64); ok && i64 >= 0 && i64 <= math.MaxUint32 {
|
||||
return any(uint32(i64)).(T), true
|
||||
}
|
||||
case uint64:
|
||||
if i64, ok := src.(int64); ok && i64 >= 0 {
|
||||
return any(uint64(i64)).(T), true
|
||||
}
|
||||
case float32:
|
||||
if f64, ok := src.(float64); ok {
|
||||
return any(float32(f64)).(T), true
|
||||
}
|
||||
case float64:
|
||||
if f64, ok := src.(float64); ok {
|
||||
return any(f64).(T), true
|
||||
}
|
||||
case string:
|
||||
switch v := src.(type) {
|
||||
case string:
|
||||
return any(v).(T), true
|
||||
case []byte:
|
||||
return any(string(v)).(T), true
|
||||
}
|
||||
case []byte:
|
||||
switch v := src.(type) {
|
||||
case []byte:
|
||||
// Copy the driver-owned buffer: database/sql may reuse it on the next row.
|
||||
return any(append([]byte(nil), v...)).(T), true
|
||||
case string:
|
||||
return any([]byte(v)).(T), true
|
||||
}
|
||||
case bool:
|
||||
if b, ok := src.(bool); ok {
|
||||
return any(b).(T), true
|
||||
}
|
||||
case time.Time:
|
||||
if t, ok := src.(time.Time); ok {
|
||||
return any(t).(T), true
|
||||
}
|
||||
}
|
||||
|
||||
return zero, false
|
||||
}
|
||||
|
||||
// convertToDriverValue safely converts primitive Go types to a value
|
||||
// compatible with database/sql driver.Value.
|
||||
//
|
||||
// Supported conversions:
|
||||
// - int, int8, int16, int32 → int64
|
||||
// - uint8, uint16, uint32 → int64
|
||||
// - uint, uint64 → int64 (only if <= math.MaxInt64)
|
||||
// - float32 → float64
|
||||
//
|
||||
// Returns the converted value and true on success, otherwise nil and false.
|
||||
func convertToDriverValue(val any) (driver.Value, bool) {
|
||||
switch v := val.(type) {
|
||||
case int:
|
||||
return int64(v), true
|
||||
case int8:
|
||||
return int64(v), true
|
||||
case int16:
|
||||
return int64(v), true
|
||||
case int32:
|
||||
return int64(v), true
|
||||
case uint:
|
||||
if uint64(v) <= math.MaxInt64 {
|
||||
return int64(v), true
|
||||
}
|
||||
case uint8:
|
||||
return int64(v), true
|
||||
case uint16:
|
||||
return int64(v), true
|
||||
case uint32:
|
||||
return int64(v), true
|
||||
case uint64:
|
||||
if v <= math.MaxInt64 {
|
||||
return int64(v), true
|
||||
}
|
||||
case float32:
|
||||
return float64(v), true
|
||||
}
|
||||
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// fitsInt returns true if i fits in a Go int.
|
||||
func fitsInt(i int64) bool { return i >= math.MinInt && i <= math.MaxInt }
|
||||
|
||||
// fitsInt8 returns true if i fits in an int8.
|
||||
func fitsInt8(i int64) bool { return i >= math.MinInt8 && i <= math.MaxInt8 }
|
||||
|
||||
// fitsInt16 returns true if i fits in an int16.
|
||||
func fitsInt16(i int64) bool { return i >= math.MinInt16 && i <= math.MaxInt16 }
|
||||
|
||||
// fitsInt32 returns true if i fits in an int32.
|
||||
func fitsInt32(i int64) bool { return i >= math.MinInt32 && i <= math.MaxInt32 }
|
||||
+1038
File diff suppressed because it is too large.
Load diff
+98
@@ -0,0 +1,98 @@
|
||||
package rand
|
||||
|
||||
import (
|
||||
crand "crypto/rand"
|
||||
"encoding/binary"
|
||||
|
||||
"github.com/enetx/g"
|
||||
)
|
||||
|
||||
// SecureBytes returns length cryptographically secure random bytes drawn from
|
||||
// crypto/rand. A zero or negative length yields nil.
|
||||
//
|
||||
// Unlike the rest of this package, the result is safe for keys, tokens and
|
||||
// other security-sensitive material.
|
||||
func SecureBytes(length g.Int) g.Bytes {
|
||||
if length <= 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
buf := make(g.Bytes, length)
|
||||
if _, err := crand.Read(buf); err != nil {
|
||||
panic(err) // crypto/rand.Read does not fail on supported platforms
|
||||
}
|
||||
|
||||
return buf
|
||||
}
|
||||
|
||||
// SecureString generates a cryptographically secure random String of the
|
||||
// specified length, selecting characters from predefined sets. If additional
|
||||
// character sets are provided, only those are used; the default set
|
||||
// (g.ASCIILetters and g.Digits) is excluded unless explicitly provided.
|
||||
//
|
||||
// If length is zero or negative, an empty String is returned. If an explicit
|
||||
// letter set is provided but resolves to empty, an empty String is returned as
|
||||
// well.
|
||||
//
|
||||
// Characters are drawn from crypto/rand with rejection sampling, so the
|
||||
// selection is uniform (no modulo bias). Unlike [String], the result is safe
|
||||
// for tokens, one-time codes and other security-sensitive material.
|
||||
//
|
||||
// rand.SecureString(32) // 32 alphanumeric characters
|
||||
// rand.SecureString(6, g.Digits) // 6-digit one-time code
|
||||
func SecureString(length g.Int, letters ...g.String) g.String {
|
||||
if length <= 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
var chars []rune
|
||||
if len(letters) != 0 {
|
||||
var buf g.Builder
|
||||
for _, set := range letters {
|
||||
_, _ = buf.WriteString(set)
|
||||
}
|
||||
|
||||
chars = buf.String().Runes()
|
||||
} else {
|
||||
chars = (g.ASCIILetters + g.Digits).Runes()
|
||||
}
|
||||
|
||||
n := len(chars)
|
||||
if n == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
var b g.Builder
|
||||
b.Grow(length)
|
||||
|
||||
for range length {
|
||||
b.WriteRune(chars[secureIndex(n)])
|
||||
}
|
||||
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// secureIndex returns a uniform random index in [0, n) drawn from crypto/rand,
|
||||
// using rejection sampling to avoid modulo bias.
|
||||
func secureIndex(n int) int {
|
||||
if n == 1 {
|
||||
return 0
|
||||
}
|
||||
|
||||
// Values at or above limit fall into the biased tail and are redrawn.
|
||||
// A limit of zero means 2^32 is an exact multiple of n — nothing to reject.
|
||||
limit := uint32((1 << 32 / uint64(n)) * uint64(n))
|
||||
|
||||
var buf [4]byte
|
||||
|
||||
for {
|
||||
if _, err := crand.Read(buf[:]); err != nil {
|
||||
panic(err) // crypto/rand.Read does not fail on supported platforms
|
||||
}
|
||||
|
||||
v := binary.BigEndian.Uint32(buf[:])
|
||||
if limit == 0 || v < limit {
|
||||
return int(v % uint32(n))
|
||||
}
|
||||
}
|
||||
}
|
||||
+204
@@ -0,0 +1,204 @@
|
||||
package g
|
||||
|
||||
import (
|
||||
"encoding/json/jsontext"
|
||||
json "encoding/json/v2"
|
||||
"errors"
|
||||
)
|
||||
|
||||
// errResultJSONShape is returned when a Result document is not a JSON object
|
||||
// with exactly one of the keys "ok" or "err".
|
||||
var errResultJSONShape = errors.New("g.Result: expected a JSON object with exactly one of the keys \"ok\" or \"err\"")
|
||||
|
||||
// MarshalJSON implements the json.Marshaler interface (encoding/json v1) for Result[T].
|
||||
// The encoding is externally tagged:
|
||||
// Ok(value) is marshaled as {"ok": <json of value>} and Err(err) is marshaled
|
||||
// as {"err": "<err.Error()>"}.
|
||||
//
|
||||
// If marshaling the contained value fails, that error is returned.
|
||||
//
|
||||
// BREAKING: the implementation is backed by encoding/json/v2, which changes some
|
||||
// edge-case semantics compared to the previous encoding/json implementation:
|
||||
// - Ok of a nil slice marshals as {"ok":[]} and Ok of a nil map as {"ok":{}},
|
||||
// rather than {"ok":null}.
|
||||
// - Strings containing invalid UTF-8 (in the Ok value or the error message)
|
||||
// are rejected with an error instead of being silently replaced with U+FFFD.
|
||||
//
|
||||
// NOTE: only the error message is serialized. The concrete error type and any
|
||||
// wrapped errors are lost — after a round trip, errors.Is/errors.As chains no
|
||||
// longer match.
|
||||
func (r Result[T]) MarshalJSON() ([]byte, error) {
|
||||
if r.IsErr() {
|
||||
msg, err := json.Marshal(r.err.Error())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return append(append([]byte(`{"err":`), msg...), '}'), nil
|
||||
}
|
||||
|
||||
v, err := json.Marshal(r.v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return append(append([]byte(`{"ok":`), v...), '}'), nil
|
||||
}
|
||||
|
||||
// UnmarshalJSON implements the json.Unmarshaler interface (encoding/json v1) for Result[T].
|
||||
// It expects the externally tagged encoding produced by MarshalJSON: a JSON
|
||||
// object with exactly one of the keys "ok" or "err".
|
||||
//
|
||||
// BREAKING: the implementation is backed by encoding/json/v2, so duplicate keys
|
||||
// are rejected — a document that repeats the same key ({"ok":1,"ok":2}) is now
|
||||
// an unmarshal error instead of the previous encoding/json last-wins semantics.
|
||||
//
|
||||
// {"err": "msg"} is unmarshaled as Err(errors.New("msg")); the value must be
|
||||
// a JSON string. {"ok": <v>} is unmarshaled as Ok with v decoded into T —
|
||||
// {"ok": null} decodes null into T following the encoding/json/v2 rules
|
||||
// (zero/nil). Anything else (both keys, neither key, extra keys, duplicate
|
||||
// keys, a non-object, or JSON null) is an unmarshal error.
|
||||
//
|
||||
// NOTE: only the error message survives a round trip. The original error type
|
||||
// is not restored — the decoded error is a plain errors.New value, so
|
||||
// errors.Is/errors.As chains against the original error no longer match.
|
||||
func (r *Result[T]) UnmarshalJSON(data []byte) error {
|
||||
var raw map[string]jsontext.Value
|
||||
if err := json.Unmarshal(data, &raw); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// JSON null decodes into a nil map without error; reject it explicitly.
|
||||
if raw == nil || len(raw) != 1 {
|
||||
return errResultJSONShape
|
||||
}
|
||||
|
||||
if msg, ok := raw["err"]; ok {
|
||||
if string(msg) == "null" {
|
||||
return errors.New("g.Result: \"err\" value must be a JSON string, got null")
|
||||
}
|
||||
|
||||
var s string
|
||||
if err := json.Unmarshal(msg, &s); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
*r = Err[T](errors.New(s))
|
||||
return nil
|
||||
}
|
||||
|
||||
value, ok := raw["ok"]
|
||||
if !ok {
|
||||
return errResultJSONShape
|
||||
}
|
||||
|
||||
var v T
|
||||
if err := json.Unmarshal(value, &v); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
*r = Ok(v)
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarshalJSONTo implements the json.MarshalerTo interface (encoding/json/v2) for Result[T].
|
||||
// encoding/json/v2 prefers this method over MarshalJSON.
|
||||
// The encoding matches MarshalJSON: Ok(value) is encoded as {"ok": <json of value>}
|
||||
// and Err(err) is encoded as {"err": "<err.Error()>"}.
|
||||
func (r Result[T]) MarshalJSONTo(enc *jsontext.Encoder) error {
|
||||
if err := enc.WriteToken(jsontext.BeginObject); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if r.IsErr() {
|
||||
if err := enc.WriteToken(jsontext.String("err")); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := enc.WriteToken(jsontext.String(r.err.Error())); err != nil {
|
||||
return err
|
||||
}
|
||||
} else {
|
||||
if err := enc.WriteToken(jsontext.String("ok")); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := json.MarshalEncode(enc, r.v); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return enc.WriteToken(jsontext.EndObject)
|
||||
}
|
||||
|
||||
// UnmarshalJSONFrom implements the json.UnmarshalerFrom interface (encoding/json/v2) for Result[T].
|
||||
// encoding/json/v2 prefers this method over UnmarshalJSON.
|
||||
// It expects a JSON object with exactly one member whose key is "ok" or "err",
|
||||
// as produced by MarshalJSONTo.
|
||||
//
|
||||
// Duplicate keys are rejected: the strict single-member contract refuses any
|
||||
// second object member (a duplicate key included) before the v2 decoder's own
|
||||
// duplicate-name check even fires, and the encoding/json/v2 decoder itself
|
||||
// forbids duplicate object member names by default. Presence of both distinct
|
||||
// keys, extra keys, neither key, a non-object, or JSON null is likewise an
|
||||
// unmarshal error.
|
||||
func (r *Result[T]) UnmarshalJSONFrom(dec *jsontext.Decoder) error {
|
||||
if dec.PeekKind() != '{' {
|
||||
// Surface the underlying syntax error for malformed input; otherwise
|
||||
// report the shape violation (null, array, string, number, bool).
|
||||
if _, err := dec.ReadToken(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return errResultJSONShape
|
||||
}
|
||||
|
||||
if _, err := dec.ReadToken(); err != nil { // consume '{'
|
||||
return err
|
||||
}
|
||||
|
||||
if dec.PeekKind() == '}' {
|
||||
return errResultJSONShape // empty object
|
||||
}
|
||||
|
||||
name, err := dec.ReadToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var res Result[T]
|
||||
|
||||
switch name.String() {
|
||||
case "err":
|
||||
if dec.PeekKind() != '"' {
|
||||
return errors.New("g.Result: \"err\" value must be a JSON string")
|
||||
}
|
||||
|
||||
var s string
|
||||
if err := json.UnmarshalDecode(dec, &s); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
res = Err[T](errors.New(s))
|
||||
case "ok":
|
||||
var v T
|
||||
if err := json.UnmarshalDecode(dec, &v); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
res = Ok(v)
|
||||
default:
|
||||
return errResultJSONShape
|
||||
}
|
||||
|
||||
if dec.PeekKind() != '}' {
|
||||
return errResultJSONShape // a second member: both keys, extra or duplicate keys
|
||||
}
|
||||
|
||||
if _, err := dec.ReadToken(); err != nil { // consume '}'
|
||||
return err
|
||||
}
|
||||
|
||||
*r = res
|
||||
return nil
|
||||
}
|
||||
+119
@@ -0,0 +1,119 @@
|
||||
package g
|
||||
|
||||
import "sync"
|
||||
|
||||
// RwLock is a reader-writer lock that protects a value of type T.
|
||||
// It allows multiple readers or a single writer at any point in time.
|
||||
// Unlike sync.RWMutex, it binds the protected data to the lock itself,
|
||||
// making it impossible to access the data without holding the lock.
|
||||
//
|
||||
// An RwLock must not be copied after first use: it embeds a sync.RWMutex and a
|
||||
// copy would protect a different value than the original (go vet's copylocks
|
||||
// analyzer flags such copies). Always pass an *RwLock, never an RwLock by value.
|
||||
type RwLock[T any] struct {
|
||||
mu sync.RWMutex
|
||||
val T
|
||||
}
|
||||
|
||||
// RwLockReadGuard provides read-only access to the value protected by an RwLock.
|
||||
// Multiple read guards can exist simultaneously.
|
||||
//
|
||||
// A guard holds pointers into its owning RwLock; copying the guard and calling
|
||||
// Unlock on more than one copy releases the same read lock twice, which panics.
|
||||
// Use each guard exactly once and do not copy it.
|
||||
type RwLockReadGuard[T any] struct {
|
||||
mu *sync.RWMutex
|
||||
val *T
|
||||
}
|
||||
|
||||
// RwLockWriteGuard provides exclusive read-write access to the value protected by an RwLock.
|
||||
// Only one write guard can exist at a time, and no read guards can coexist with it.
|
||||
//
|
||||
// A guard holds pointers into its owning RwLock; copying the guard and calling
|
||||
// Unlock on more than one copy releases the same write lock twice, which panics.
|
||||
// Use each guard exactly once and do not copy it.
|
||||
type RwLockWriteGuard[T any] struct {
|
||||
mu *sync.RWMutex
|
||||
val *T
|
||||
}
|
||||
|
||||
// NewRwLock creates a new RwLock containing the given value.
|
||||
func NewRwLock[T any](value T) *RwLock[T] { return &RwLock[T]{val: value} }
|
||||
|
||||
// RWith acquires a read lock, calls fn with a copy of the protected value,
|
||||
// and releases the lock when fn returns.
|
||||
// The value is passed by copy to prevent accidental mutation under a read lock.
|
||||
func (r *RwLock[T]) RWith(fn func(T)) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
fn(r.val)
|
||||
}
|
||||
|
||||
// With acquires a write lock, calls fn with a pointer to the protected value,
|
||||
// and releases the lock when fn returns.
|
||||
// This is a convenience method that eliminates the need for manual Lock/Unlock management.
|
||||
func (r *RwLock[T]) With(fn func(*T)) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
fn(&r.val)
|
||||
}
|
||||
|
||||
// Read acquires a read lock and returns a guard that provides read-only access.
|
||||
// Multiple goroutines can hold read locks simultaneously.
|
||||
// The caller must call Unlock on the guard when done (typically via defer).
|
||||
func (r *RwLock[T]) Read() RwLockReadGuard[T] {
|
||||
r.mu.RLock()
|
||||
return RwLockReadGuard[T]{mu: &r.mu, val: &r.val}
|
||||
}
|
||||
|
||||
// Write acquires a write lock and returns a guard that provides exclusive access.
|
||||
// No other readers or writers can access the value while this guard exists.
|
||||
// The caller must call Unlock on the guard when done (typically via defer).
|
||||
func (r *RwLock[T]) Write() RwLockWriteGuard[T] {
|
||||
r.mu.Lock()
|
||||
return RwLockWriteGuard[T]{mu: &r.mu, val: &r.val}
|
||||
}
|
||||
|
||||
// TryRead attempts to acquire a read lock without blocking.
|
||||
// Returns Some(guard) if successful, None if a write lock is held.
|
||||
func (r *RwLock[T]) TryRead() Option[RwLockReadGuard[T]] {
|
||||
if r.mu.TryRLock() {
|
||||
return Some(RwLockReadGuard[T]{mu: &r.mu, val: &r.val})
|
||||
}
|
||||
|
||||
return None[RwLockReadGuard[T]]()
|
||||
}
|
||||
|
||||
// TryWrite attempts to acquire a write lock without blocking.
|
||||
// Returns Some(guard) if successful, None if any lock is held.
|
||||
func (r *RwLock[T]) TryWrite() Option[RwLockWriteGuard[T]] {
|
||||
if r.mu.TryLock() {
|
||||
return Some(RwLockWriteGuard[T]{mu: &r.mu, val: &r.val})
|
||||
}
|
||||
|
||||
return None[RwLockWriteGuard[T]]()
|
||||
}
|
||||
|
||||
// Get returns a copy of the protected value.
|
||||
func (g RwLockReadGuard[T]) Get() T { return *g.val }
|
||||
|
||||
// Deref returns a pointer to the protected value for direct access.
|
||||
// Note: modifying through this pointer would be a logic error.
|
||||
func (g RwLockReadGuard[T]) Deref() *T { return g.val }
|
||||
|
||||
// Unlock releases the read lock. Must be called when done with the guard.
|
||||
func (g RwLockReadGuard[T]) Unlock() { g.mu.RUnlock() }
|
||||
|
||||
// Get returns a copy of the protected value.
|
||||
func (g RwLockWriteGuard[T]) Get() T { return *g.val }
|
||||
|
||||
// Set replaces the protected value with a new one.
|
||||
func (g RwLockWriteGuard[T]) Set(value T) { *g.val = value }
|
||||
|
||||
// Deref returns a pointer to the protected value for direct manipulation.
|
||||
func (g RwLockWriteGuard[T]) Deref() *T { return g.val }
|
||||
|
||||
// Unlock releases the write lock. Must be called when done with the guard.
|
||||
func (g RwLockWriteGuard[T]) Unlock() { g.mu.Unlock() }
|
||||
+1959
File diff suppressed because it is too large.
Load diff
+1152
File diff suppressed because it is too large.
Load diff
+147
@@ -0,0 +1,147 @@
|
||||
package g
|
||||
|
||||
import "github.com/enetx/g/cmp"
|
||||
|
||||
// Collect returns a collector over the sequence. Collect itself is lazy and
|
||||
// does not consume the sequence; only the collector's materializer methods do,
|
||||
// turning the elements into a concrete container:
|
||||
//
|
||||
// s.Iter().Filter(fn).Collect().Slice()
|
||||
// s.Iter().Filter(fn).Collect().Set[int]()
|
||||
// s.Iter().Collect().Heap(cmp.Cmp)
|
||||
// s.Iter().Collect().Deque()
|
||||
//
|
||||
// Set (like Map/MapOrd/MapSafe on the key-value collector) takes the element
|
||||
// type as an explicit type argument: Go checks constraints at method
|
||||
// DECLARATION, where the sequence's own type parameter is still `any`, so a
|
||||
// method of Seq[V any] can never mention Set[V] — the comparable constraint
|
||||
// must live on the method's own type parameter, and that parameter cannot be
|
||||
// inferred from zero arguments. The other collectors need no annotations.
|
||||
func (seq Seq[V]) Collect() collector[V] { return collector[V]{seq} }
|
||||
|
||||
// Collect returns a collector over the key-value sequence. Collect itself is
|
||||
// lazy and does not consume the sequence; only the collector's materializer
|
||||
// methods do, turning the pairs into a concrete container:
|
||||
//
|
||||
// m.Iter().FilterByKey(fn).Collect().Pairs()
|
||||
// m.Iter().FilterByKey(fn).Collect().Map[string, string]()
|
||||
// mo.Iter().Collect().MapOrd[string, string]()
|
||||
func (seq Seq2[K, V]) Collect() collector2[K, V] { return collector2[K, V]{seq} }
|
||||
|
||||
// collector materializes a value sequence into containers; build one with
|
||||
// Seq.Collect.
|
||||
type collector[V any] struct{ seq Seq[V] }
|
||||
|
||||
// collector2 materializes a key-value sequence into containers; build one with
|
||||
// Seq2.Collect.
|
||||
type collector2[K, V any] struct{ seq Seq2[K, V] }
|
||||
|
||||
// Slice consumes the sequence and returns its elements as a Slice.
|
||||
func (c collector[V]) Slice() Slice[V] { return c.seq.seqToSlice() }
|
||||
|
||||
// Deque consumes the sequence and returns its elements as a Deque, preserving
|
||||
// order.
|
||||
func (c collector[V]) Deque() *Deque[V] {
|
||||
result := NewDeque[V]()
|
||||
|
||||
c.seq(func(v V) bool {
|
||||
result.PushBack(v)
|
||||
return true
|
||||
})
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// Heap consumes the sequence and returns its elements as a Heap ordered by
|
||||
// compareFn.
|
||||
func (c collector[V]) Heap(compareFn func(V, V) cmp.Ordering) *Heap[V] {
|
||||
result := NewHeap(compareFn)
|
||||
|
||||
c.seq(func(v V) bool {
|
||||
result.Push(v)
|
||||
return true
|
||||
})
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// Set consumes the sequence and returns its elements as a Set, deduplicating
|
||||
// them. The element type is passed explicitly — Collect().Set[int]() — because
|
||||
// the sequence's own type parameter cannot carry the comparable constraint;
|
||||
// each element is converted at runtime and a mismatched type argument panics.
|
||||
func (c collector[V]) Set[W comparable]() Set[W] {
|
||||
collection := make(Set[W])
|
||||
|
||||
c.seq(func(v V) bool {
|
||||
collection[any(v).(W)] = Unit{}
|
||||
return true
|
||||
})
|
||||
|
||||
return collection
|
||||
}
|
||||
|
||||
// Pairs consumes the sequence and returns its elements as plain pairs.
|
||||
func (c collector2[K, V]) Pairs() []Pair[K, V] {
|
||||
var result []Pair[K, V]
|
||||
|
||||
c.seq(func(k K, v V) bool {
|
||||
result = append(result, Pair[K, V]{Key: k, Value: v})
|
||||
return true
|
||||
})
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// Map consumes the key-value sequence and returns it as a Map; later keys
|
||||
// overwrite earlier ones. Both types are passed explicitly, mirroring the
|
||||
// container's own signature — Collect().Map[string, int]() — because the
|
||||
// sequence's key parameter cannot carry the comparable constraint; the pairs
|
||||
// are converted at runtime and a mismatched type argument panics.
|
||||
func (c collector2[K, V]) Map[K2 comparable, V2 any]() Map[K2, V2] {
|
||||
collection := NewMap[K2, V2]()
|
||||
|
||||
c.seq(func(k K, v V) bool {
|
||||
collection[any(k).(K2)] = any(v).(V2)
|
||||
return true
|
||||
})
|
||||
|
||||
return collection
|
||||
}
|
||||
|
||||
// MapSafe consumes the key-value sequence and returns it as a thread-safe
|
||||
// MapSafe; later keys overwrite earlier ones. Both types are passed explicitly
|
||||
// — Collect().MapSafe[string, int]() — see [collector2.Map].
|
||||
func (c collector2[K, V]) MapSafe[K2 comparable, V2 any]() *MapSafe[K2, V2] {
|
||||
collection := NewMapSafe[K2, V2]()
|
||||
|
||||
c.seq(func(k K, v V) bool {
|
||||
collection.Insert(any(k).(K2), any(v).(V2))
|
||||
return true
|
||||
})
|
||||
|
||||
return collection
|
||||
}
|
||||
|
||||
// MapOrd consumes the key-value sequence and returns it as a MapOrd, keeping
|
||||
// first-seen key order; a repeated key updates the value in place. Both types
|
||||
// are passed explicitly — Collect().MapOrd[string, int]() — see
|
||||
// [collector2.Map].
|
||||
func (c collector2[K, V]) MapOrd[K2 comparable, V2 any]() MapOrd[K2, V2] {
|
||||
collection := NewMapOrd[K2, V2]()
|
||||
idx := make(map[K2]int)
|
||||
|
||||
c.seq(func(k K, v V) bool {
|
||||
key := any(k).(K2)
|
||||
if i, ok := idx[key]; ok {
|
||||
collection[i].Value = any(v).(V2)
|
||||
return true
|
||||
}
|
||||
|
||||
collection = append(collection, Pair[K2, V2]{Key: key, Value: any(v).(V2)})
|
||||
idx[key] = len(collection) - 1
|
||||
|
||||
return true
|
||||
})
|
||||
|
||||
return collection
|
||||
}
|
||||
+456
@@ -0,0 +1,456 @@
|
||||
// Copyright 2025 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package http
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"github.com/enetx/http/httptrace"
|
||||
"net/url"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// A ClientConn is a client connection to an HTTP server.
|
||||
//
|
||||
// Unlike a [Transport], a ClientConn represents a single connection.
|
||||
// Most users should use a Transport rather than creating client connections directly.
|
||||
type ClientConn struct {
|
||||
cc genericClientConn
|
||||
|
||||
stateHookMu sync.Mutex
|
||||
userStateHook func(*ClientConn)
|
||||
stateHookRunning bool
|
||||
lastAvailable int
|
||||
lastInFlight int
|
||||
lastClosed bool
|
||||
}
|
||||
|
||||
// newClientConner is the interface implemented by HTTP/2 transports to create new client conns.
|
||||
//
|
||||
// The http package (this package) needs a way to ask the http2 package to
|
||||
// create a client connection.
|
||||
//
|
||||
// Transport.TLSNextProto["h2"] contains a function which appears to do this,
|
||||
// but for historical reasons it does not: The TLSNextProto function adds a
|
||||
// *tls.Conn to the http2.Transport's connection pool and returns a RoundTripper
|
||||
// which is backed by that connection pool. NewClientConn needs a way to get a
|
||||
// single client connection out of the http2 package.
|
||||
//
|
||||
// The http2 package registers a RoundTripper with Transport.RegisterProtocol.
|
||||
// If this RoundTripper implements newClientConner, then Transport.NewClientConn will use
|
||||
// it to create new HTTP/2 client connections.
|
||||
type newClientConner interface {
|
||||
// NewClientConn creates a new client connection from a net.Conn.
|
||||
//
|
||||
// The RoundTripper returned by NewClientConn must implement genericClientConn.
|
||||
// (We don't define NewClientConn as returning genericClientConn,
|
||||
// because either we'd need to make genericClientConn an exported type
|
||||
// or define it as a type alias. Neither is particularly appealing.)
|
||||
//
|
||||
// The state hook passed here is the internal state hook
|
||||
// (ClientConn.maybeRunStateHook). The internal state hook calls
|
||||
// the user state hook (if any), which is set by the user with
|
||||
// ClientConn.SetStateHook.
|
||||
//
|
||||
// The client connection should arrange to call the internal state hook
|
||||
// when the connection closes, when requests complete, and when the
|
||||
// connection concurrency limit changes.
|
||||
//
|
||||
// The client connection must call the internal state hook when the connection state
|
||||
// changes asynchronously, such as when a request completes.
|
||||
//
|
||||
// The internal state hook need not be called after synchronous changes to the state:
|
||||
// Close, Reserve, Release, and RoundTrip calls which don't start a request
|
||||
// do not need to call the hook.
|
||||
//
|
||||
// The general idea is that if we call (for example) Close,
|
||||
// we know that the connection state has probably changed and we
|
||||
// don't need the state hook to tell us that.
|
||||
// However, if the connection closes asynchronously
|
||||
// (because, for example, the other end of the conn closed it),
|
||||
// the state hook needs to inform us.
|
||||
NewClientConn(nc net.Conn, internalStateHook func()) (RoundTripper, error)
|
||||
}
|
||||
|
||||
// genericClientConn is an interface implemented by HTTP/2 client conns
|
||||
// returned from newClientConner.NewClientConn.
|
||||
//
|
||||
// See the newClientConner doc comment for more information.
|
||||
type genericClientConn interface {
|
||||
Close() error
|
||||
Err() error
|
||||
RoundTrip(req *Request) (*Response, error)
|
||||
Reserve() error
|
||||
Release()
|
||||
Available() int
|
||||
InFlight() int
|
||||
}
|
||||
|
||||
// NewClientConn creates a new client connection to the given address.
|
||||
//
|
||||
// If scheme is "http", the connection is unencrypted.
|
||||
// If scheme is "https", the connection uses TLS.
|
||||
//
|
||||
// The protocol used for the new connection is determined by the scheme,
|
||||
// Transport.Protocols configuration field, and protocols supported by the
|
||||
// server. See Transport.Protocols for more details.
|
||||
//
|
||||
// If Transport.Proxy is set and indicates that a request sent to the given
|
||||
// address should use a proxy, the new connection uses that proxy.
|
||||
//
|
||||
// NewClientConn always creates a new connection,
|
||||
// even if the Transport has an existing cached connection to the given host.
|
||||
//
|
||||
// The new connection is not added to the Transport's connection cache,
|
||||
// and will not be used by [Transport.RoundTrip].
|
||||
// It does not count against the MaxIdleConns and MaxConnsPerHost limits.
|
||||
//
|
||||
// The caller is responsible for closing the new connection.
|
||||
func (t *Transport) NewClientConn(ctx context.Context, scheme, address string) (*ClientConn, error) {
|
||||
t.nextProtoOnce.Do(t.onceSetNextProtoDefaults)
|
||||
|
||||
switch scheme {
|
||||
case "http", "https":
|
||||
default:
|
||||
return nil, fmt.Errorf("net/http: invalid scheme %q", scheme)
|
||||
}
|
||||
|
||||
host, port, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if port == "" {
|
||||
port = schemePort(scheme)
|
||||
}
|
||||
|
||||
var proxyURL *url.URL
|
||||
if t.Proxy != nil {
|
||||
// Transport.Proxy takes a *Request, so create a fake one to pass it.
|
||||
req := &Request{
|
||||
ctx: ctx,
|
||||
Method: "GET",
|
||||
URL: &url.URL{
|
||||
Scheme: scheme,
|
||||
Host: host,
|
||||
Path: "/",
|
||||
},
|
||||
Proto: "HTTP/1.1",
|
||||
ProtoMajor: 1,
|
||||
ProtoMinor: 1,
|
||||
Header: make(Header),
|
||||
Body: NoBody,
|
||||
Host: host,
|
||||
}
|
||||
var err error
|
||||
proxyURL, err = t.Proxy(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
cm := connectMethod{
|
||||
targetScheme: scheme,
|
||||
targetAddr: net.JoinHostPort(host, port),
|
||||
proxyURL: proxyURL,
|
||||
}
|
||||
|
||||
// The state hook is a bit tricky:
|
||||
// The persistConn has a state hook which calls ClientConn.maybeRunStateHook,
|
||||
// which in turn calls the user-provided state hook (if any).
|
||||
//
|
||||
// ClientConn.maybeRunStateHook handles debouncing hook calls for both
|
||||
// HTTP/1 and HTTP/2.
|
||||
//
|
||||
// Since there's no need to change the persistConn's hook, we set it at creation time.
|
||||
cc := &ClientConn{}
|
||||
const isClientConn = true
|
||||
pconn, err := t.dialConn(ctx, cm, isClientConn, cc.maybeRunStateHook)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Note that cc.maybeRunStateHook may have been called
|
||||
// in the short window between dialConn and now.
|
||||
// This is fine.
|
||||
cc.stateHookMu.Lock()
|
||||
defer cc.stateHookMu.Unlock()
|
||||
if pconn.alt != nil {
|
||||
// If pconn.alt is set, this is a connection implemented in another package
|
||||
// (probably x/net/http2) or the bundled copy in h2_bundle.go.
|
||||
gc, ok := pconn.alt.(genericClientConn)
|
||||
if !ok {
|
||||
return nil, errors.New("http: NewClientConn returned something that is not a ClientConn")
|
||||
}
|
||||
cc.cc = gc
|
||||
cc.lastAvailable = gc.Available()
|
||||
} else {
|
||||
// This is an HTTP/1 connection.
|
||||
pconn.availch = make(chan struct{}, 1)
|
||||
pconn.availch <- struct{}{}
|
||||
cc.cc = http1ClientConn{pconn}
|
||||
cc.lastAvailable = 1
|
||||
}
|
||||
return cc, nil
|
||||
}
|
||||
|
||||
// Close closes the connection.
|
||||
// Outstanding RoundTrip calls are interrupted.
|
||||
func (cc *ClientConn) Close() error {
|
||||
defer cc.maybeRunStateHook()
|
||||
return cc.cc.Close()
|
||||
}
|
||||
|
||||
// Err reports any fatal connection errors.
|
||||
// It returns nil if the connection is usable.
|
||||
// If it returns non-nil, the connection can no longer be used.
|
||||
func (cc *ClientConn) Err() error {
|
||||
return cc.cc.Err()
|
||||
}
|
||||
|
||||
func validateClientConnRequest(req *Request) error {
|
||||
if req.URL == nil {
|
||||
return errors.New("http: nil Request.URL")
|
||||
}
|
||||
if req.Header == nil {
|
||||
return errors.New("http: nil Request.Header")
|
||||
}
|
||||
// Validate the outgoing headers.
|
||||
if err := validateHeaders(req.Header); err != "" {
|
||||
return fmt.Errorf("http: invalid header %s", err)
|
||||
}
|
||||
// Validate the outgoing trailers too.
|
||||
if err := validateHeaders(req.Trailer); err != "" {
|
||||
return fmt.Errorf("http: invalid trailer %s", err)
|
||||
}
|
||||
if req.Method != "" && !validMethod(req.Method) {
|
||||
return fmt.Errorf("http: invalid method %q", req.Method)
|
||||
}
|
||||
if req.URL.Host == "" {
|
||||
return errors.New("http: no Host in request URL")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RoundTrip implements the [RoundTripper] interface.
|
||||
//
|
||||
// The request is sent on the client connection,
|
||||
// regardless of the URL being requested or any proxy settings.
|
||||
//
|
||||
// If the connection is at its concurrency limit,
|
||||
// RoundTrip waits for the connection to become available
|
||||
// before sending the request.
|
||||
func (cc *ClientConn) RoundTrip(req *Request) (*Response, error) {
|
||||
defer cc.maybeRunStateHook()
|
||||
if err := validateClientConnRequest(req); err != nil {
|
||||
cc.Release()
|
||||
return nil, err
|
||||
}
|
||||
return cc.cc.RoundTrip(req)
|
||||
}
|
||||
|
||||
// Available reports the number of requests that may be sent
|
||||
// to the connection without blocking.
|
||||
// It returns 0 if the connection is closed.
|
||||
func (cc *ClientConn) Available() int {
|
||||
return cc.cc.Available()
|
||||
}
|
||||
|
||||
// InFlight reports the number of requests in flight,
|
||||
// including reserved requests.
|
||||
// It returns 0 if the connection is closed.
|
||||
func (cc *ClientConn) InFlight() int {
|
||||
return cc.cc.InFlight()
|
||||
}
|
||||
|
||||
// Reserve reserves a concurrency slot on the connection.
|
||||
// If Reserve returns nil, one additional RoundTrip call may be made
|
||||
// without waiting for an existing request to complete.
|
||||
//
|
||||
// The reserved concurrency slot is accounted as an in-flight request.
|
||||
// A successful call to RoundTrip will decrement the Available count
|
||||
// and increment the InFlight count.
|
||||
//
|
||||
// Each successful call to Reserve should be followed by exactly one call
|
||||
// to RoundTrip or Release, which will consume or release the reservation.
|
||||
//
|
||||
// If the connection is closed or at its concurrency limit,
|
||||
// Reserve returns an error.
|
||||
func (cc *ClientConn) Reserve() error {
|
||||
defer cc.maybeRunStateHook()
|
||||
return cc.cc.Reserve()
|
||||
}
|
||||
|
||||
// Release releases an unused concurrency slot reserved by Reserve.
|
||||
// If there are no reserved concurrency slots, it has no effect.
|
||||
func (cc *ClientConn) Release() {
|
||||
defer cc.maybeRunStateHook()
|
||||
cc.cc.Release()
|
||||
}
|
||||
|
||||
// shouldRunStateHook returns the user's state hook if we should call it,
|
||||
// or nil if we don't need to call it at this time.
|
||||
func (cc *ClientConn) shouldRunStateHook(stopRunning bool) func(*ClientConn) {
|
||||
cc.stateHookMu.Lock()
|
||||
defer cc.stateHookMu.Unlock()
|
||||
if cc.cc == nil {
|
||||
return nil
|
||||
}
|
||||
if stopRunning {
|
||||
cc.stateHookRunning = false
|
||||
}
|
||||
if cc.userStateHook == nil {
|
||||
return nil
|
||||
}
|
||||
if cc.stateHookRunning {
|
||||
return nil
|
||||
}
|
||||
var (
|
||||
available = cc.Available()
|
||||
inFlight = cc.InFlight()
|
||||
closed = cc.Err() != nil
|
||||
)
|
||||
var hook func(*ClientConn)
|
||||
if available > cc.lastAvailable || inFlight < cc.lastInFlight || closed != cc.lastClosed {
|
||||
hook = cc.userStateHook
|
||||
cc.stateHookRunning = true
|
||||
}
|
||||
cc.lastAvailable = available
|
||||
cc.lastInFlight = inFlight
|
||||
cc.lastClosed = closed
|
||||
return hook
|
||||
}
|
||||
|
||||
func (cc *ClientConn) maybeRunStateHook() {
|
||||
hook := cc.shouldRunStateHook(false)
|
||||
if hook == nil {
|
||||
return
|
||||
}
|
||||
// Run the hook synchronously.
|
||||
//
|
||||
// This means that if, for example, the user calls resp.Body.Close to finish a request,
|
||||
// the Close call will synchronously run the hook, giving the hook the chance to
|
||||
// return the ClientConn to a connection pool before the next request is made.
|
||||
hook(cc)
|
||||
// The connection state may have changed while the hook was running,
|
||||
// in which case we need to run it again.
|
||||
//
|
||||
// If we do need to run the hook again, do so in a new goroutine to avoid blocking
|
||||
// the current goroutine indefinitely.
|
||||
hook = cc.shouldRunStateHook(true)
|
||||
if hook != nil {
|
||||
go func() {
|
||||
for hook != nil {
|
||||
hook(cc)
|
||||
hook = cc.shouldRunStateHook(true)
|
||||
}
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
// SetStateHook arranges for f to be called when the state of the connection changes.
|
||||
// At most one call to f is made at a time.
|
||||
// If the connection's state has changed since it was created,
|
||||
// f is called immediately in a separate goroutine.
|
||||
// f may be called synchronously from RoundTrip or Response.Body.Close.
|
||||
//
|
||||
// If SetStateHook is called multiple times, the new hook replaces the old one.
|
||||
// If f is nil, no further calls will be made to f after SetStateHook returns.
|
||||
//
|
||||
// f is called when Available increases (more requests may be sent on the connection),
|
||||
// InFlight decreases (existing requests complete), or Err begins returning non-nil
|
||||
// (the connection is no longer usable).
|
||||
func (cc *ClientConn) SetStateHook(f func(*ClientConn)) {
|
||||
cc.stateHookMu.Lock()
|
||||
cc.userStateHook = f
|
||||
cc.stateHookMu.Unlock()
|
||||
cc.maybeRunStateHook()
|
||||
}
|
||||
|
||||
// http1ClientConn is a genericClientConn implementation backed by
|
||||
// an HTTP/1 *persistConn (pconn.alt is nil).
|
||||
type http1ClientConn struct {
|
||||
pconn *persistConn
|
||||
}
|
||||
|
||||
func (cc http1ClientConn) RoundTrip(req *Request) (*Response, error) {
|
||||
ctx := req.Context()
|
||||
trace := httptrace.ContextClientTrace(ctx)
|
||||
|
||||
// Convert Request.Cancel into context cancelation.
|
||||
ctx, cancel := context.WithCancelCause(req.Context())
|
||||
if req.Cancel != nil {
|
||||
go awaitLegacyCancel(ctx, cancel, req)
|
||||
}
|
||||
|
||||
treq := &transportRequest{Request: req, trace: trace, ctx: ctx, cancel: cancel}
|
||||
resp, err := cc.pconn.roundTrip(treq)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp.Request = req
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
func (cc http1ClientConn) Close() error {
|
||||
cc.pconn.close(errors.New("ClientConn closed"))
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cc http1ClientConn) Err() error {
|
||||
select {
|
||||
case <-cc.pconn.closech:
|
||||
return cc.pconn.closed
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (cc http1ClientConn) Available() int {
|
||||
cc.pconn.mu.Lock()
|
||||
defer cc.pconn.mu.Unlock()
|
||||
if cc.pconn.closed != nil || cc.pconn.reserved || cc.pconn.inFlight {
|
||||
return 0
|
||||
}
|
||||
return 1
|
||||
}
|
||||
|
||||
func (cc http1ClientConn) InFlight() int {
|
||||
cc.pconn.mu.Lock()
|
||||
defer cc.pconn.mu.Unlock()
|
||||
if cc.pconn.closed == nil && (cc.pconn.reserved || cc.pconn.inFlight) {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (cc http1ClientConn) Reserve() error {
|
||||
cc.pconn.mu.Lock()
|
||||
defer cc.pconn.mu.Unlock()
|
||||
if cc.pconn.closed != nil {
|
||||
return cc.pconn.closed
|
||||
}
|
||||
select {
|
||||
case <-cc.pconn.availch:
|
||||
default:
|
||||
return errors.New("connection is unavailable")
|
||||
}
|
||||
cc.pconn.reserved = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cc http1ClientConn) Release() {
|
||||
cc.pconn.mu.Lock()
|
||||
defer cc.pconn.mu.Unlock()
|
||||
if cc.pconn.reserved {
|
||||
select {
|
||||
case cc.pconn.availch <- struct{}{}:
|
||||
default:
|
||||
panic("cannot release reservation")
|
||||
}
|
||||
cc.pconn.reserved = false
|
||||
}
|
||||
}
|
||||
+20
@@ -0,0 +1,20 @@
|
||||
// Copyright 2026 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
//go:build !go1.27
|
||||
|
||||
package http2
|
||||
|
||||
import "github.com/enetx/http"
|
||||
|
||||
// Support for go.dev/issue/75500 is added in Go 1.27. In case anyone uses
|
||||
// x/net with versions before Go 1.27, we return true here so that their write
|
||||
// scheduler will still be the round-robin write scheduler rather than the RFC
|
||||
// 9218 write scheduler. That way, older users of Go will not see a sudden
|
||||
// change of behavior just from importing x/net.
|
||||
//
|
||||
// TODO(nsh): remove this file after x/net go.mod is at Go 1.27.
|
||||
func clientPriorityDisabled(_ *http.Server) bool {
|
||||
return true
|
||||
}
|
||||
+13
@@ -0,0 +1,13 @@
|
||||
// Copyright 2026 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
//go:build go1.27
|
||||
|
||||
package http2
|
||||
|
||||
import "github.com/enetx/http"
|
||||
|
||||
func clientPriorityDisabled(s *http.Server) bool {
|
||||
return s.DisableClientPriority
|
||||
}
|
||||
+665
@@ -0,0 +1,665 @@
|
||||
// Copyright 2025 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
// Package httpsfv provides functionality for dealing with HTTP Structured
|
||||
// Field Values.
|
||||
package httpsfv
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
func isLCAlpha(b byte) bool {
|
||||
return (b >= 'a' && b <= 'z')
|
||||
}
|
||||
|
||||
func isAlpha(b byte) bool {
|
||||
return isLCAlpha(b) || (b >= 'A' && b <= 'Z')
|
||||
}
|
||||
|
||||
func isDigit(b byte) bool {
|
||||
return b >= '0' && b <= '9'
|
||||
}
|
||||
|
||||
func isVChar(b byte) bool {
|
||||
return b >= 0x21 && b <= 0x7e
|
||||
}
|
||||
|
||||
func isSP(b byte) bool {
|
||||
return b == 0x20
|
||||
}
|
||||
|
||||
func isTChar(b byte) bool {
|
||||
if isAlpha(b) || isDigit(b) {
|
||||
return true
|
||||
}
|
||||
return slices.Contains([]byte{'!', '#', '$', '%', '&', '\'', '*', '+', '-', '.', '^', '_', '`', '|', '~'}, b)
|
||||
}
|
||||
|
||||
func countLeftWhitespace(s string) int {
|
||||
i := 0
|
||||
for _, ch := range []byte(s) {
|
||||
if ch != ' ' && ch != '\t' {
|
||||
break
|
||||
}
|
||||
i++
|
||||
}
|
||||
return i
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc4648#section-8.
|
||||
func decOctetHex(ch1, ch2 byte) (ch byte, ok bool) {
|
||||
decBase16 := func(in byte) (out byte, ok bool) {
|
||||
if !isDigit(in) && !(in >= 'a' && in <= 'f') {
|
||||
return 0, false
|
||||
}
|
||||
if isDigit(in) {
|
||||
return in - '0', true
|
||||
}
|
||||
return in - 'a' + 10, true
|
||||
}
|
||||
|
||||
if ch1, ok = decBase16(ch1); !ok {
|
||||
return 0, ok
|
||||
}
|
||||
if ch2, ok = decBase16(ch2); !ok {
|
||||
return 0, ok
|
||||
}
|
||||
return ch1<<4 | ch2, true
|
||||
}
|
||||
|
||||
// ParseList parses a list from a given HTTP Structured Field Values.
|
||||
//
|
||||
// Given an HTTP SFV string that represents a list, it will call the given
|
||||
// function using each of the members and parameters contained in the list.
|
||||
// This allows the caller to extract information out of the list.
|
||||
//
|
||||
// This function will return once it encounters the end of the string, or
|
||||
// something that is not a list. If it cannot consume the entire given
|
||||
// string, the ok value returned will be false.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-list.
|
||||
func ParseList(s string, f func(member, param string)) (ok bool) {
|
||||
for len(s) != 0 {
|
||||
var member, param string
|
||||
if len(s) != 0 && s[0] == '(' {
|
||||
if member, s, ok = consumeBareInnerList(s, nil); !ok {
|
||||
return ok
|
||||
}
|
||||
} else {
|
||||
if member, s, ok = consumeBareItem(s); !ok {
|
||||
return ok
|
||||
}
|
||||
}
|
||||
if param, s, ok = consumeParameter(s, nil); !ok {
|
||||
return ok
|
||||
}
|
||||
if f != nil {
|
||||
f(member, param)
|
||||
}
|
||||
|
||||
s = s[countLeftWhitespace(s):]
|
||||
if len(s) == 0 {
|
||||
break
|
||||
}
|
||||
if s[0] != ',' {
|
||||
return false
|
||||
}
|
||||
s = s[1:]
|
||||
s = s[countLeftWhitespace(s):]
|
||||
if len(s) == 0 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// consumeBareInnerList consumes an inner list
|
||||
// (https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-an-inner-list),
|
||||
// except for the inner list's top-most parameter.
|
||||
// For example, given `(a;b c;d);e`, it will consume only `(a;b c;d)`.
|
||||
func consumeBareInnerList(s string, f func(bareItem, param string)) (consumed, rest string, ok bool) {
|
||||
if len(s) == 0 || s[0] != '(' {
|
||||
return "", s, false
|
||||
}
|
||||
rest = s[1:]
|
||||
for len(rest) != 0 {
|
||||
var bareItem, param string
|
||||
rest = rest[countLeftWhitespace(rest):]
|
||||
if len(rest) != 0 && rest[0] == ')' {
|
||||
rest = rest[1:]
|
||||
break
|
||||
}
|
||||
if bareItem, rest, ok = consumeBareItem(rest); !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
if param, rest, ok = consumeParameter(rest, nil); !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
if len(rest) == 0 || (rest[0] != ')' && !isSP(rest[0])) {
|
||||
return "", s, false
|
||||
}
|
||||
if f != nil {
|
||||
f(bareItem, param)
|
||||
}
|
||||
}
|
||||
return s[:len(s)-len(rest)], rest, true
|
||||
}
|
||||
|
||||
// ParseBareInnerList parses a bare inner list from a given HTTP Structured
|
||||
// Field Values.
|
||||
//
|
||||
// We define a bare inner list as an inner list
|
||||
// (https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-an-inner-list),
|
||||
// without the top-most parameter of the inner list. For example, given the
|
||||
// inner list `(a;b c;d);e`, the bare inner list would be `(a;b c;d)`.
|
||||
//
|
||||
// Given an HTTP SFV string that represents a bare inner list, it will call the
|
||||
// given function using each of the bare item and parameter within the bare
|
||||
// inner list. This allows the caller to extract information out of the bare
|
||||
// inner list.
|
||||
//
|
||||
// This function will return once it encounters the end of the bare inner list,
|
||||
// or something that is not a bare inner list. If it cannot consume the entire
|
||||
// given string, the ok value returned will be false.
|
||||
func ParseBareInnerList(s string, f func(bareItem, param string)) (ok bool) {
|
||||
_, rest, ok := consumeBareInnerList(s, f)
|
||||
return rest == "" && ok
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-an-item.
|
||||
func consumeItem(s string, f func(bareItem, param string)) (consumed, rest string, ok bool) {
|
||||
var bareItem, param string
|
||||
if bareItem, rest, ok = consumeBareItem(s); !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
if param, rest, ok = consumeParameter(rest, nil); !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
if f != nil {
|
||||
f(bareItem, param)
|
||||
}
|
||||
return s[:len(s)-len(rest)], rest, true
|
||||
}
|
||||
|
||||
// ParseItem parses an item from a given HTTP Structured Field Values.
|
||||
//
|
||||
// Given an HTTP SFV string that represents an item, it will call the given
|
||||
// function once, with the bare item and the parameter of the item. This allows
|
||||
// the caller to extract information out of the item.
|
||||
//
|
||||
// This function will return once it encounters the end of the string, or
|
||||
// something that is not an item. If it cannot consume the entire given
|
||||
// string, the ok value returned will be false.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-an-item.
|
||||
func ParseItem(s string, f func(bareItem, param string)) (ok bool) {
|
||||
_, rest, ok := consumeItem(s, f)
|
||||
return rest == "" && ok
|
||||
}
|
||||
|
||||
// ParseDictionary parses a dictionary from a given HTTP Structured Field
|
||||
// Values.
|
||||
//
|
||||
// Given an HTTP SFV string that represents a dictionary, it will call the
|
||||
// given function using each of the keys, values, and parameters contained in
|
||||
// the dictionary. This allows the caller to extract information out of the
|
||||
// dictionary.
|
||||
//
|
||||
// This function will return once it encounters the end of the string, or
|
||||
// something that is not a dictionary. If it cannot consume the entire given
|
||||
// string, the ok value returned will be false.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-dictionary.
|
||||
func ParseDictionary(s string, f func(key, val, param string)) (ok bool) {
|
||||
for len(s) != 0 {
|
||||
var key, val, param string
|
||||
val = "?1" // Default value for empty val is boolean true.
|
||||
if key, s, ok = consumeKey(s); !ok {
|
||||
return ok
|
||||
}
|
||||
if len(s) != 0 && s[0] == '=' {
|
||||
s = s[1:]
|
||||
if len(s) != 0 && s[0] == '(' {
|
||||
if val, s, ok = consumeBareInnerList(s, nil); !ok {
|
||||
return ok
|
||||
}
|
||||
} else {
|
||||
if val, s, ok = consumeBareItem(s); !ok {
|
||||
return ok
|
||||
}
|
||||
}
|
||||
}
|
||||
if param, s, ok = consumeParameter(s, nil); !ok {
|
||||
return ok
|
||||
}
|
||||
if f != nil {
|
||||
f(key, val, param)
|
||||
}
|
||||
s = s[countLeftWhitespace(s):]
|
||||
if len(s) == 0 {
|
||||
break
|
||||
}
|
||||
if s[0] == ',' {
|
||||
s = s[1:]
|
||||
}
|
||||
s = s[countLeftWhitespace(s):]
|
||||
if len(s) == 0 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#parse-param.
|
||||
func consumeParameter(s string, f func(key, val string)) (consumed, rest string, ok bool) {
|
||||
rest = s
|
||||
for len(rest) != 0 {
|
||||
var key, val string
|
||||
val = "?1" // Default value for empty val is boolean true.
|
||||
if rest[0] != ';' {
|
||||
break
|
||||
}
|
||||
rest = rest[1:]
|
||||
rest = rest[countLeftWhitespace(rest):]
|
||||
key, rest, ok = consumeKey(rest)
|
||||
if !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
if len(rest) != 0 && rest[0] == '=' {
|
||||
rest = rest[1:]
|
||||
val, rest, ok = consumeBareItem(rest)
|
||||
if !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
}
|
||||
if f != nil {
|
||||
f(key, val)
|
||||
}
|
||||
}
|
||||
return s[:len(s)-len(rest)], rest, true
|
||||
}
|
||||
|
||||
// ParseParameter parses a parameter from a given HTTP Structured Field Values.
|
||||
//
|
||||
// Given an HTTP SFV string that represents a parameter, it will call the given
|
||||
// function using each of the keys and values contained in the parameter. This
|
||||
// allows the caller to extract information out of the parameter.
|
||||
//
|
||||
// This function will return once it encounters the end of the string, or
|
||||
// something that is not a parameter. If it cannot consume the entire given
|
||||
// string, the ok value returned will be false.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#parse-param.
|
||||
func ParseParameter(s string, f func(key, val string)) (ok bool) {
|
||||
_, rest, ok := consumeParameter(s, f)
|
||||
return rest == "" && ok
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-key.
|
||||
func consumeKey(s string) (consumed, rest string, ok bool) {
|
||||
if len(s) == 0 || (!isLCAlpha(s[0]) && s[0] != '*') {
|
||||
return "", s, false
|
||||
}
|
||||
i := 0
|
||||
for _, ch := range []byte(s) {
|
||||
if !isLCAlpha(ch) && !isDigit(ch) && !slices.Contains([]byte("_-.*"), ch) {
|
||||
break
|
||||
}
|
||||
i++
|
||||
}
|
||||
return s[:i], s[i:], true
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-an-integer-or-decim.
|
||||
func consumeIntegerOrDecimal(s string) (consumed, rest string, ok bool) {
|
||||
var i, signOffset, periodIndex int
|
||||
var isDecimal bool
|
||||
if i < len(s) && s[i] == '-' {
|
||||
i++
|
||||
signOffset++
|
||||
}
|
||||
if i >= len(s) {
|
||||
return "", s, false
|
||||
}
|
||||
if !isDigit(s[i]) {
|
||||
return "", s, false
|
||||
}
|
||||
for i < len(s) {
|
||||
ch := s[i]
|
||||
if isDigit(ch) {
|
||||
i++
|
||||
continue
|
||||
}
|
||||
if !isDecimal && ch == '.' {
|
||||
if i-signOffset > 12 {
|
||||
return "", s, false
|
||||
}
|
||||
periodIndex = i
|
||||
isDecimal = true
|
||||
i++
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
if !isDecimal && i-signOffset > 15 {
|
||||
return "", s, false
|
||||
}
|
||||
if isDecimal {
|
||||
if i-signOffset > 16 {
|
||||
return "", s, false
|
||||
}
|
||||
if s[i-1] == '.' {
|
||||
return "", s, false
|
||||
}
|
||||
if i-periodIndex-1 > 3 {
|
||||
return "", s, false
|
||||
}
|
||||
}
|
||||
return s[:i], s[i:], true
|
||||
}
|
||||
|
||||
// ParseInteger parses an integer from a given HTTP Structured Field Values.
|
||||
//
|
||||
// The entire HTTP SFV string must consist of a valid integer. It returns the
|
||||
// parsed integer and an ok boolean value, indicating success or not.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-an-integer-or-decim.
|
||||
func ParseInteger(s string) (parsed int64, ok bool) {
|
||||
if _, rest, ok := consumeIntegerOrDecimal(s); !ok || rest != "" {
|
||||
return 0, false
|
||||
}
|
||||
if n, err := strconv.ParseInt(s, 10, 64); err == nil {
|
||||
return n, true
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
// ParseDecimal parses a decimal from a given HTTP Structured Field Values.
|
||||
//
|
||||
// The entire HTTP SFV string must consist of a valid decimal. It returns the
|
||||
// parsed decimal and an ok boolean value, indicating success or not.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-an-integer-or-decim.
|
||||
func ParseDecimal(s string) (parsed float64, ok bool) {
|
||||
if _, rest, ok := consumeIntegerOrDecimal(s); !ok || rest != "" {
|
||||
return 0, false
|
||||
}
|
||||
if !strings.Contains(s, ".") {
|
||||
return 0, false
|
||||
}
|
||||
if n, err := strconv.ParseFloat(s, 64); err == nil {
|
||||
return n, true
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-string.
|
||||
func consumeString(s string) (consumed, rest string, ok bool) {
|
||||
if len(s) == 0 || s[0] != '"' {
|
||||
return "", s, false
|
||||
}
|
||||
for i := 1; i < len(s); i++ {
|
||||
switch ch := s[i]; ch {
|
||||
case '\\':
|
||||
if i+1 >= len(s) {
|
||||
return "", s, false
|
||||
}
|
||||
i++
|
||||
if ch = s[i]; ch != '"' && ch != '\\' {
|
||||
return "", s, false
|
||||
}
|
||||
case '"':
|
||||
return s[:i+1], s[i+1:], true
|
||||
default:
|
||||
if !isVChar(ch) && !isSP(ch) {
|
||||
return "", s, false
|
||||
}
|
||||
}
|
||||
}
|
||||
return "", s, false
|
||||
}
|
||||
|
||||
// ParseString parses a Go string from a given HTTP Structured Field Values.
|
||||
//
|
||||
// The entire HTTP SFV string must consist of a valid string. It returns the
|
||||
// parsed string and an ok boolean value, indicating success or not.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-string.
|
||||
func ParseString(s string) (parsed string, ok bool) {
|
||||
if _, rest, ok := consumeString(s); !ok || rest != "" {
|
||||
return "", false
|
||||
}
|
||||
return s[1 : len(s)-1], true
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-token
|
||||
func consumeToken(s string) (consumed, rest string, ok bool) {
|
||||
if len(s) == 0 || (!isAlpha(s[0]) && s[0] != '*') {
|
||||
return "", s, false
|
||||
}
|
||||
i := 0
|
||||
for _, ch := range []byte(s) {
|
||||
if !isTChar(ch) && !slices.Contains([]byte(":/"), ch) {
|
||||
break
|
||||
}
|
||||
i++
|
||||
}
|
||||
return s[:i], s[i:], true
|
||||
}
|
||||
|
||||
// ParseToken parses a token from a given HTTP Structured Field Values.
|
||||
//
|
||||
// The entire HTTP SFV string must consist of a valid token. It returns the
|
||||
// parsed token and an ok boolean value, indicating success or not.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-token
|
||||
func ParseToken(s string) (parsed string, ok bool) {
|
||||
if _, rest, ok := consumeToken(s); !ok || rest != "" {
|
||||
return "", false
|
||||
}
|
||||
return s, true
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-byte-sequence.
|
||||
func consumeByteSequence(s string) (consumed, rest string, ok bool) {
|
||||
if len(s) == 0 || s[0] != ':' {
|
||||
return "", s, false
|
||||
}
|
||||
for i := 1; i < len(s); i++ {
|
||||
if ch := s[i]; ch == ':' {
|
||||
return s[:i+1], s[i+1:], true
|
||||
}
|
||||
if ch := s[i]; !isAlpha(ch) && !isDigit(ch) && !slices.Contains([]byte("+/="), ch) {
|
||||
return "", s, false
|
||||
}
|
||||
}
|
||||
return "", s, false
|
||||
}
|
||||
|
||||
// ParseByteSequence parses a byte sequence from a given HTTP Structured Field
|
||||
// Values.
|
||||
//
|
||||
// The entire HTTP SFV string must consist of a valid byte sequence. It returns
|
||||
// the parsed byte sequence and an ok boolean value, indicating success or not.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-byte-sequence.
|
||||
func ParseByteSequence(s string) (parsed []byte, ok bool) {
|
||||
if _, rest, ok := consumeByteSequence(s); !ok || rest != "" {
|
||||
return nil, false
|
||||
}
|
||||
return []byte(s[1 : len(s)-1]), true
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-boolean.
|
||||
func consumeBoolean(s string) (consumed, rest string, ok bool) {
|
||||
if len(s) >= 2 && (s[:2] == "?0" || s[:2] == "?1") {
|
||||
return s[:2], s[2:], true
|
||||
}
|
||||
return "", s, false
|
||||
}
|
||||
|
||||
// ParseBoolean parses a boolean from a given HTTP Structured Field Values.
|
||||
//
|
||||
// The entire HTTP SFV string must consist of a valid boolean. It returns the
|
||||
// parsed boolean and an ok boolean value, indicating success or not.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-boolean.
|
||||
func ParseBoolean(s string) (parsed bool, ok bool) {
|
||||
if _, rest, ok := consumeBoolean(s); !ok || rest != "" {
|
||||
return false, false
|
||||
}
|
||||
return s == "?1", true
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-date.
|
||||
func consumeDate(s string) (consumed, rest string, ok bool) {
|
||||
if len(s) == 0 || s[0] != '@' {
|
||||
return "", s, false
|
||||
}
|
||||
if _, rest, ok = consumeIntegerOrDecimal(s[1:]); !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
consumed = s[:len(s)-len(rest)]
|
||||
if slices.Contains([]byte(consumed), '.') {
|
||||
return "", s, false
|
||||
}
|
||||
return consumed, rest, ok
|
||||
}
|
||||
|
||||
// ParseDate parses a date from a given HTTP Structured Field Values.
|
||||
//
|
||||
// The entire HTTP SFV string must consist of a valid date. It returns the
|
||||
// parsed date and an ok boolean value, indicating success or not.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-date.
|
||||
func ParseDate(s string) (parsed time.Time, ok bool) {
|
||||
if _, rest, ok := consumeDate(s); !ok || rest != "" {
|
||||
return time.Time{}, false
|
||||
}
|
||||
if n, ok := ParseInteger(s[1:]); !ok {
|
||||
return time.Time{}, false
|
||||
} else {
|
||||
return time.Unix(n, 0), true
|
||||
}
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-display-string.
|
||||
func consumeDisplayString(s string) (consumed, rest string, ok bool) {
|
||||
// To prevent excessive allocation, especially when input is large, we
|
||||
// maintain a buffer of 4 bytes to keep track of the last rune we
|
||||
// encounter. This way, we can validate that the display string conforms to
|
||||
// UTF-8 without actually building the whole string.
|
||||
var lastRune [4]byte
|
||||
var runeLen int
|
||||
isPartOfValidRune := func(ch byte) bool {
|
||||
lastRune[runeLen] = ch
|
||||
runeLen++
|
||||
if utf8.FullRune(lastRune[:runeLen]) {
|
||||
r, s := utf8.DecodeRune(lastRune[:runeLen])
|
||||
if r == utf8.RuneError {
|
||||
return false
|
||||
}
|
||||
copy(lastRune[:], lastRune[s:runeLen])
|
||||
runeLen -= s
|
||||
return true
|
||||
}
|
||||
return runeLen <= 4
|
||||
}
|
||||
|
||||
if len(s) <= 1 || s[:2] != `%"` {
|
||||
return "", s, false
|
||||
}
|
||||
i := 2
|
||||
for i < len(s) {
|
||||
ch := s[i]
|
||||
if !isVChar(ch) && !isSP(ch) {
|
||||
return "", s, false
|
||||
}
|
||||
switch ch {
|
||||
case '"':
|
||||
if runeLen > 0 {
|
||||
return "", s, false
|
||||
}
|
||||
return s[:i+1], s[i+1:], true
|
||||
case '%':
|
||||
if i+2 >= len(s) {
|
||||
return "", s, false
|
||||
}
|
||||
if ch, ok = decOctetHex(s[i+1], s[i+2]); !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
if ok = isPartOfValidRune(ch); !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
i += 3
|
||||
default:
|
||||
if ok = isPartOfValidRune(ch); !ok {
|
||||
return "", s, ok
|
||||
}
|
||||
i++
|
||||
}
|
||||
}
|
||||
return "", s, false
|
||||
}
|
||||
|
||||
// ParseDisplayString parses a display string from a given HTTP Structured
|
||||
// Field Values.
|
||||
//
|
||||
// The entire HTTP SFV string must consist of a valid display string. It
|
||||
// returns the parsed display string and an ok boolean value, indicating
|
||||
// success or not.
|
||||
//
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#name-parsing-a-display-string.
|
||||
func ParseDisplayString(s string) (parsed string, ok bool) {
|
||||
if _, rest, ok := consumeDisplayString(s); !ok || rest != "" {
|
||||
return "", false
|
||||
}
|
||||
// consumeDisplayString() already validates that we have a valid display
|
||||
// string. Therefore, we can just construct the display string, without
|
||||
// validating it again.
|
||||
s = s[2 : len(s)-1]
|
||||
var b strings.Builder
|
||||
for i := 0; i < len(s); {
|
||||
if s[i] == '%' {
|
||||
decoded, _ := decOctetHex(s[i+1], s[i+2])
|
||||
b.WriteByte(decoded)
|
||||
i += 3
|
||||
continue
|
||||
}
|
||||
b.WriteByte(s[i])
|
||||
i++
|
||||
}
|
||||
return b.String(), true
|
||||
}
|
||||
|
||||
// https://www.rfc-editor.org/rfc/rfc9651.html#parse-bare-item.
|
||||
func consumeBareItem(s string) (consumed, rest string, ok bool) {
|
||||
if len(s) == 0 {
|
||||
return "", s, false
|
||||
}
|
||||
ch := s[0]
|
||||
switch {
|
||||
case ch == '-' || isDigit(ch):
|
||||
return consumeIntegerOrDecimal(s)
|
||||
case ch == '"':
|
||||
return consumeString(s)
|
||||
case ch == '*' || isAlpha(ch):
|
||||
return consumeToken(s)
|
||||
case ch == ':':
|
||||
return consumeByteSequence(s)
|
||||
case ch == '?':
|
||||
return consumeBoolean(s)
|
||||
case ch == '@':
|
||||
return consumeDate(s)
|
||||
case ch == '%':
|
||||
return consumeDisplayString(s)
|
||||
default:
|
||||
return "", s, false
|
||||
}
|
||||
}
|
||||
+224
@@ -0,0 +1,224 @@
|
||||
// Copyright 2025 The Go Authors. All rights reserved.
|
||||
// Use of this source code is governed by a BSD-style
|
||||
// license that can be found in the LICENSE file.
|
||||
|
||||
package http2
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math"
|
||||
)
|
||||
|
||||
type streamMetadata struct {
|
||||
location *writeQueue
|
||||
priority PriorityParam
|
||||
}
|
||||
|
||||
type priorityWriteSchedulerRFC9218 struct {
|
||||
// control contains control frames (SETTINGS, PING, etc.).
|
||||
control writeQueue
|
||||
|
||||
// heads contain the head of a circular list of streams.
|
||||
// We put these heads within a nested array that represents urgency and
|
||||
// incremental, as defined in
|
||||
// https://www.rfc-editor.org/rfc/rfc9218.html#name-priority-parameters.
|
||||
// 8 represents u=0 up to u=7, and 2 represents i=false and i=true.
|
||||
heads [8][2]*writeQueue
|
||||
|
||||
// streams contains a mapping between each stream ID and their metadata, so
|
||||
// we can quickly locate them when needing to, for example, adjust their
|
||||
// priority.
|
||||
streams map[uint32]streamMetadata
|
||||
|
||||
// queuePool are empty queues for reuse.
|
||||
queuePool writeQueuePool
|
||||
|
||||
// prioritizeIncremental is used to determine whether we should prioritize
|
||||
// incremental streams or not, when urgency is the same in a given Pop()
|
||||
// call.
|
||||
prioritizeIncremental bool
|
||||
|
||||
// priorityUpdateBuf is used to buffer the most recent PRIORITY_UPDATE we
|
||||
// receive per https://www.rfc-editor.org/rfc/rfc9218.html#name-the-priority_update-frame.
|
||||
priorityUpdateBuf struct {
|
||||
// streamID being 0 means that the buffer is empty. This is a safe
|
||||
// assumption as PRIORITY_UPDATE for stream 0 is a PROTOCOL_ERROR.
|
||||
streamID uint32
|
||||
priority PriorityParam
|
||||
}
|
||||
}
|
||||
|
||||
func newPriorityWriteSchedulerRFC9218() WriteScheduler {
|
||||
ws := &priorityWriteSchedulerRFC9218{
|
||||
streams: make(map[uint32]streamMetadata),
|
||||
}
|
||||
return ws
|
||||
}
|
||||
|
||||
func (ws *priorityWriteSchedulerRFC9218) OpenStream(streamID uint32, opt OpenStreamOptions) {
|
||||
if ws.streams[streamID].location != nil {
|
||||
panic(fmt.Errorf("stream %d already opened", streamID))
|
||||
}
|
||||
if streamID == ws.priorityUpdateBuf.streamID {
|
||||
ws.priorityUpdateBuf.streamID = 0
|
||||
opt.priority = ws.priorityUpdateBuf.priority
|
||||
}
|
||||
q := ws.queuePool.get()
|
||||
ws.streams[streamID] = streamMetadata{
|
||||
location: q,
|
||||
priority: opt.priority,
|
||||
}
|
||||
|
||||
u, i := opt.priority.urgency, opt.priority.incremental
|
||||
if ws.heads[u][i] == nil {
|
||||
ws.heads[u][i] = q
|
||||
q.next = q
|
||||
q.prev = q
|
||||
} else {
|
||||
// Queues are stored in a ring.
|
||||
// Insert the new stream before ws.head, putting it at the end of the list.
|
||||
q.prev = ws.heads[u][i].prev
|
||||
q.next = ws.heads[u][i]
|
||||
q.prev.next = q
|
||||
q.next.prev = q
|
||||
}
|
||||
}
|
||||
|
||||
func (ws *priorityWriteSchedulerRFC9218) CloseStream(streamID uint32) {
|
||||
metadata := ws.streams[streamID]
|
||||
q, u, i := metadata.location, metadata.priority.urgency, metadata.priority.incremental
|
||||
if q == nil {
|
||||
return
|
||||
}
|
||||
if q.next == q {
|
||||
// This was the only open stream.
|
||||
ws.heads[u][i] = nil
|
||||
} else {
|
||||
q.prev.next = q.next
|
||||
q.next.prev = q.prev
|
||||
if ws.heads[u][i] == q {
|
||||
ws.heads[u][i] = q.next
|
||||
}
|
||||
}
|
||||
delete(ws.streams, streamID)
|
||||
ws.queuePool.put(q)
|
||||
}
|
||||
|
||||
func (ws *priorityWriteSchedulerRFC9218) AdjustStream(streamID uint32, priority PriorityParam) {
|
||||
metadata := ws.streams[streamID]
|
||||
q, u, i := metadata.location, metadata.priority.urgency, metadata.priority.incremental
|
||||
if q == nil {
|
||||
ws.priorityUpdateBuf.streamID = streamID
|
||||
ws.priorityUpdateBuf.priority = priority
|
||||
return
|
||||
}
|
||||
|
||||
// Remove stream from current location.
|
||||
if q.next == q {
|
||||
// This was the only open stream.
|
||||
ws.heads[u][i] = nil
|
||||
} else {
|
||||
q.prev.next = q.next
|
||||
q.next.prev = q.prev
|
||||
if ws.heads[u][i] == q {
|
||||
ws.heads[u][i] = q.next
|
||||
}
|
||||
}
|
||||
|
||||
// Insert stream to the new queue.
|
||||
u, i = priority.urgency, priority.incremental
|
||||
if ws.heads[u][i] == nil {
|
||||
ws.heads[u][i] = q
|
||||
q.next = q
|
||||
q.prev = q
|
||||
} else {
|
||||
// Queues are stored in a ring.
|
||||
// Insert the new stream before ws.head, putting it at the end of the list.
|
||||
q.prev = ws.heads[u][i].prev
|
||||
q.next = ws.heads[u][i]
|
||||
q.prev.next = q
|
||||
q.next.prev = q
|
||||
}
|
||||
|
||||
// Update the metadata.
|
||||
ws.streams[streamID] = streamMetadata{
|
||||
location: q,
|
||||
priority: priority,
|
||||
}
|
||||
}
|
||||
|
||||
func (ws *priorityWriteSchedulerRFC9218) Push(wr FrameWriteRequest) {
|
||||
if wr.isControl() {
|
||||
ws.control.push(wr)
|
||||
return
|
||||
}
|
||||
q := ws.streams[wr.StreamID()].location
|
||||
if q == nil {
|
||||
// This is a closed stream.
|
||||
// wr should not be a HEADERS or DATA frame.
|
||||
// We push the request onto the control queue.
|
||||
if wr.DataSize() > 0 {
|
||||
panic("add DATA on non-open stream")
|
||||
}
|
||||
ws.control.push(wr)
|
||||
return
|
||||
}
|
||||
q.push(wr)
|
||||
}
|
||||
|
||||
func (ws *priorityWriteSchedulerRFC9218) Pop() (FrameWriteRequest, bool) {
|
||||
// Control and RST_STREAM frames first.
|
||||
if !ws.control.empty() {
|
||||
return ws.control.shift(), true
|
||||
}
|
||||
|
||||
// On the next Pop(), we want to prioritize incremental if we prioritized
|
||||
// non-incremental request of the same urgency this time. Vice-versa.
|
||||
// i.e. when there are incremental and non-incremental requests at the same
|
||||
// priority, we give 50% of our bandwidth to the incremental ones in
|
||||
// aggregate and 50% to the first non-incremental one (since
|
||||
// non-incremental streams do not use round-robin writes).
|
||||
ws.prioritizeIncremental = !ws.prioritizeIncremental
|
||||
|
||||
// Always prioritize lowest u (i.e. highest urgency level).
|
||||
for u := range ws.heads {
|
||||
for i := range ws.heads[u] {
|
||||
// When we want to prioritize incremental, we try to pop i=true
|
||||
// first before i=false when u is the same.
|
||||
if ws.prioritizeIncremental {
|
||||
i = (i + 1) % 2
|
||||
}
|
||||
q := ws.heads[u][i]
|
||||
if q == nil {
|
||||
continue
|
||||
}
|
||||
for {
|
||||
if wr, ok := q.consume(math.MaxInt32); ok {
|
||||
if i == 1 {
|
||||
// For incremental streams, we update head to q.next so
|
||||
// we can round-robin between multiple streams that can
|
||||
// immediately benefit from partial writes.
|
||||
ws.heads[u][i] = q.next
|
||||
} else {
|
||||
// For non-incremental streams, we try to finish one to
|
||||
// completion rather than doing round-robin. However,
|
||||
// we update head here so that if q.consume() is !ok
|
||||
// (e.g. the stream has no more frame to consume), head
|
||||
// is updated to the next q that has frames to consume
|
||||
// on future iterations. This way, we do not prioritize
|
||||
// writing to unavailable stream on next Pop() calls,
|
||||
// preventing head-of-line blocking.
|
||||
ws.heads[u][i] = q
|
||||
}
|
||||
return wr, true
|
||||
}
|
||||
q = q.next
|
||||
if q == ws.heads[u][i] {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
return FrameWriteRequest{}, false
|
||||
}
|
||||
+9
@@ -0,0 +1,9 @@
|
||||
# HTTP/3
|
||||
|
||||
[](https://quic-go.net/docs/)
|
||||
[](https://pkg.go.dev/github.com/quic-go/quic-go/http3)
|
||||
|
||||
This package implements HTTP/3 ([RFC 9114](https://datatracker.ietf.org/doc/html/rfc9114)), including QPACK ([RFC 9204](https://datatracker.ietf.org/doc/html/rfc9204)) and HTTP Datagrams ([RFC 9297](https://datatracker.ietf.org/doc/html/rfc9297)).
|
||||
It aims to provide feature parity with the standard library's HTTP/1.1 and HTTP/2 implementation.
|
||||
|
||||
Detailed documentation can be found on [quic-go.net](https://quic-go.net/docs/).
|
||||
+137
@@ -0,0 +1,137 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"sync"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
)
|
||||
|
||||
// Settingser allows waiting for and retrieving the peer's HTTP/3 settings.
|
||||
type Settingser interface {
|
||||
// ReceivedSettings returns a channel that is closed once the peer's SETTINGS frame was received.
|
||||
// Settings can be obtained from the Settings method after the channel was closed.
|
||||
ReceivedSettings() <-chan struct{}
|
||||
// Settings returns the settings received on this connection.
|
||||
// It is only valid to call this function after the channel returned by ReceivedSettings was closed.
|
||||
Settings() *Settings
|
||||
}
|
||||
|
||||
var errTooMuchData = errors.New("peer sent too much data")
|
||||
|
||||
// The body is used in the requestBody (for a http.Request) and the responseBody (for a http.Response).
|
||||
type body struct {
|
||||
str *Stream
|
||||
|
||||
remainingContentLength int64
|
||||
violatedContentLength bool
|
||||
hasContentLength bool
|
||||
}
|
||||
|
||||
func newBody(str *Stream, contentLength int64) *body {
|
||||
b := &body{str: str}
|
||||
if contentLength >= 0 {
|
||||
b.hasContentLength = true
|
||||
b.remainingContentLength = contentLength
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func (r *body) StreamID() quic.StreamID { return r.str.StreamID() }
|
||||
|
||||
func (r *body) checkContentLengthViolation() error {
|
||||
if !r.hasContentLength {
|
||||
return nil
|
||||
}
|
||||
if r.remainingContentLength < 0 || r.remainingContentLength == 0 && r.str.hasMoreData() {
|
||||
if !r.violatedContentLength {
|
||||
r.str.CancelRead(quic.StreamErrorCode(ErrCodeMessageError))
|
||||
r.str.CancelWrite(quic.StreamErrorCode(ErrCodeMessageError))
|
||||
r.violatedContentLength = true
|
||||
}
|
||||
return errTooMuchData
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *body) Read(b []byte) (int, error) {
|
||||
if err := r.checkContentLengthViolation(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if r.hasContentLength {
|
||||
b = b[:min(int64(len(b)), r.remainingContentLength)]
|
||||
}
|
||||
n, err := r.str.Read(b)
|
||||
r.remainingContentLength -= int64(n)
|
||||
if err := r.checkContentLengthViolation(); err != nil {
|
||||
return n, err
|
||||
}
|
||||
return n, maybeReplaceError(err)
|
||||
}
|
||||
|
||||
func (r *body) Close() error {
|
||||
r.str.CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled))
|
||||
return nil
|
||||
}
|
||||
|
||||
type requestBody struct {
|
||||
body
|
||||
connCtx context.Context
|
||||
rcvdSettings <-chan struct{}
|
||||
getSettings func() *Settings
|
||||
}
|
||||
|
||||
var _ io.ReadCloser = &requestBody{}
|
||||
|
||||
func newRequestBody(str *Stream, contentLength int64, connCtx context.Context, rcvdSettings <-chan struct{}, getSettings func() *Settings) *requestBody {
|
||||
return &requestBody{
|
||||
body: *newBody(str, contentLength),
|
||||
connCtx: connCtx,
|
||||
rcvdSettings: rcvdSettings,
|
||||
getSettings: getSettings,
|
||||
}
|
||||
}
|
||||
|
||||
type hijackableBody struct {
|
||||
body body
|
||||
|
||||
// only set for the http.Response
|
||||
// The channel is closed when the user is done with this response:
|
||||
// either when Read() errors, or when Close() is called.
|
||||
reqDone chan<- struct{}
|
||||
reqDoneOnce sync.Once
|
||||
}
|
||||
|
||||
var _ io.ReadCloser = &hijackableBody{}
|
||||
|
||||
func newResponseBody(str *Stream, contentLength int64, done chan<- struct{}) *hijackableBody {
|
||||
return &hijackableBody{
|
||||
body: *newBody(str, contentLength),
|
||||
reqDone: done,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *hijackableBody) Read(b []byte) (int, error) {
|
||||
n, err := r.body.Read(b)
|
||||
if err != nil {
|
||||
r.requestDone()
|
||||
}
|
||||
return n, maybeReplaceError(err)
|
||||
}
|
||||
|
||||
func (r *hijackableBody) requestDone() {
|
||||
if r.reqDone != nil {
|
||||
r.reqDoneOnce.Do(func() {
|
||||
close(r.reqDone)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func (r *hijackableBody) Close() error {
|
||||
r.requestDone()
|
||||
// If the EOF was read, CancelRead() is a no-op.
|
||||
r.body.str.CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled))
|
||||
return nil
|
||||
}
|
||||
+61
@@ -0,0 +1,61 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"io"
|
||||
|
||||
"github.com/quic-go/quic-go/quicvarint"
|
||||
)
|
||||
|
||||
// CapsuleType is the type of the capsule
|
||||
type CapsuleType uint64
|
||||
|
||||
// CapsuleProtocolHeader is the header value used to advertise support for the capsule protocol
|
||||
const CapsuleProtocolHeader = "Capsule-Protocol"
|
||||
|
||||
type exactReader struct {
|
||||
R io.LimitedReader
|
||||
}
|
||||
|
||||
func (r *exactReader) Read(b []byte) (int, error) {
|
||||
n, err := r.R.Read(b)
|
||||
if err == io.EOF && r.R.N > 0 {
|
||||
return n, io.ErrUnexpectedEOF
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
// ParseCapsule parses the header of a Capsule.
|
||||
// It returns an io.Reader that can be used to read the Capsule value.
|
||||
// The Capsule value must be read entirely (i.e. until the io.EOF) before using r again.
|
||||
func ParseCapsule(r quicvarint.Reader) (CapsuleType, io.Reader, error) {
|
||||
cbr := countingByteReader{Reader: r}
|
||||
ct, err := quicvarint.Read(&cbr)
|
||||
if err != nil {
|
||||
// If an io.EOF is returned without consuming any bytes, return it unmodified.
|
||||
// Otherwise, return an io.ErrUnexpectedEOF.
|
||||
if err == io.EOF && cbr.NumRead > 0 {
|
||||
return 0, nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
return 0, nil, err
|
||||
}
|
||||
l, err := quicvarint.Read(r)
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
return 0, nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
return 0, nil, err
|
||||
}
|
||||
return CapsuleType(ct), &exactReader{R: io.LimitedReader{R: r, N: int64(l)}}, nil
|
||||
}
|
||||
|
||||
// WriteCapsule writes a capsule
|
||||
func WriteCapsule(w quicvarint.Writer, ct CapsuleType, value []byte) error {
|
||||
b := make([]byte, 0, 16)
|
||||
b = quicvarint.Append(b, uint64(ct))
|
||||
b = quicvarint.Append(b, uint64(len(value)))
|
||||
if _, err := w.Write(b); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := w.Write(value)
|
||||
return err
|
||||
}
|
||||
+508
@@ -0,0 +1,508 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/textproto"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/http"
|
||||
"github.com/enetx/http/httptrace"
|
||||
|
||||
"github.com/enetx/http3/qlog"
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
|
||||
"github.com/quic-go/qpack"
|
||||
)
|
||||
|
||||
const (
|
||||
// MethodGet0RTT allows a GET request to be sent using 0-RTT.
|
||||
// Note that 0-RTT doesn't provide replay protection and should only be used for idempotent requests.
|
||||
MethodGet0RTT = "GET_0RTT"
|
||||
// MethodHead0RTT allows a HEAD request to be sent using 0-RTT.
|
||||
// Note that 0-RTT doesn't provide replay protection and should only be used for idempotent requests.
|
||||
MethodHead0RTT = "HEAD_0RTT"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultUserAgent = "quic-go HTTP/3"
|
||||
defaultMaxResponseHeaderBytes = 10 * 1 << 20 // 10 MB
|
||||
)
|
||||
|
||||
var errGoAway = errors.New("connection in graceful shutdown")
|
||||
|
||||
type errConnUnusable struct{ e error }
|
||||
|
||||
func (e *errConnUnusable) Unwrap() error { return e.e }
|
||||
func (e *errConnUnusable) Error() string { return fmt.Sprintf("http3: conn unusable: %s", e.e.Error()) }
|
||||
|
||||
const max1xxResponses = 5 // arbitrary bound on number of informational responses
|
||||
|
||||
var defaultQuicConfig = &quic.Config{
|
||||
MaxIncomingStreams: -1, // don't allow the server to create bidirectional streams
|
||||
KeepAlivePeriod: 10 * time.Second,
|
||||
}
|
||||
|
||||
// ClientConn is an HTTP/3 client doing requests to a single remote server.
|
||||
type ClientConn struct {
|
||||
conn *quic.Conn
|
||||
rawConn *rawConn
|
||||
|
||||
decoder *qpack.Decoder
|
||||
|
||||
// Additional HTTP/3 settings.
|
||||
// It is invalid to specify any settings defined by RFC 9114 (HTTP/3) and RFC 9297 (HTTP Datagrams).
|
||||
additionalSettings g.MapOrd[uint64, uint64]
|
||||
|
||||
// maxResponseHeaderBytes specifies a limit on how many response bytes are
|
||||
// allowed in the server's response header.
|
||||
maxResponseHeaderBytes int
|
||||
|
||||
// disableCompression, if true, prevents the Transport from requesting compression with an
|
||||
// "Accept-Encoding: gzip" request header when the Request contains no existing Accept-Encoding value.
|
||||
// If the Transport requests gzip on its own and gets a gzipped response, it's transparently
|
||||
// decoded in the Response.Body.
|
||||
// However, if the user explicitly requested gzip it is not automatically uncompressed.
|
||||
disableCompression bool
|
||||
|
||||
streamMx sync.Mutex
|
||||
maxStreamID quic.StreamID // set once a GOAWAY frame is received
|
||||
lastStreamID quic.StreamID // the highest stream ID that was opened
|
||||
|
||||
qlogger qlogwriter.Recorder
|
||||
logger *slog.Logger
|
||||
|
||||
requestWriter *requestWriter
|
||||
}
|
||||
|
||||
var _ http.RoundTripper = &ClientConn{}
|
||||
|
||||
func newClientConn(
|
||||
conn *quic.Conn,
|
||||
enableDatagrams bool,
|
||||
additionalSettings g.MapOrd[uint64, uint64],
|
||||
maxResponseHeaderBytes int,
|
||||
disableCompression bool,
|
||||
logger *slog.Logger,
|
||||
) *ClientConn {
|
||||
var qlogger qlogwriter.Recorder
|
||||
if qlogTrace := conn.QlogTrace(); qlogTrace != nil && qlogTrace.SupportsSchemas(qlog.EventSchema) {
|
||||
qlogger = qlogTrace.AddProducer()
|
||||
}
|
||||
c := &ClientConn{
|
||||
conn: conn,
|
||||
additionalSettings: additionalSettings,
|
||||
disableCompression: disableCompression,
|
||||
maxStreamID: invalidStreamID,
|
||||
lastStreamID: invalidStreamID,
|
||||
logger: logger,
|
||||
qlogger: qlogger,
|
||||
decoder: qpack.NewDecoder(),
|
||||
}
|
||||
if maxResponseHeaderBytes <= 0 {
|
||||
c.maxResponseHeaderBytes = defaultMaxResponseHeaderBytes
|
||||
} else {
|
||||
c.maxResponseHeaderBytes = maxResponseHeaderBytes
|
||||
}
|
||||
c.requestWriter = newRequestWriter()
|
||||
c.rawConn = newRawConn(
|
||||
conn,
|
||||
enableDatagrams,
|
||||
c.onStreamsEmpty,
|
||||
c.handleControlStream,
|
||||
qlogger,
|
||||
c.logger,
|
||||
)
|
||||
// send the SETTINGs frame, using 0-RTT data, if possible
|
||||
go func() {
|
||||
_, err := c.rawConn.openControlStream(&settingsFrame{
|
||||
Datagram: enableDatagrams,
|
||||
Other: additionalSettings,
|
||||
// MaxFieldSectionSize: int64(c.maxResponseHeaderBytes),
|
||||
})
|
||||
if err != nil {
|
||||
if c.logger != nil {
|
||||
c.logger.Debug("setting up connection failed", "error", err)
|
||||
}
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeInternalError), "")
|
||||
return
|
||||
}
|
||||
}()
|
||||
return c
|
||||
}
|
||||
|
||||
// OpenRequestStream opens a new request stream on the HTTP/3 connection.
|
||||
func (c *ClientConn) OpenRequestStream(ctx context.Context) (*RequestStream, error) {
|
||||
return c.openRequestStream(ctx, c.requestWriter, nil, c.disableCompression, c.maxResponseHeaderBytes)
|
||||
}
|
||||
|
||||
func (c *ClientConn) openRequestStream(
|
||||
ctx context.Context,
|
||||
requestWriter *requestWriter,
|
||||
reqDone chan<- struct{},
|
||||
disableCompression bool,
|
||||
maxHeaderBytes int,
|
||||
) (*RequestStream, error) {
|
||||
c.streamMx.Lock()
|
||||
maxStreamID := c.maxStreamID
|
||||
var nextStreamID quic.StreamID
|
||||
if c.lastStreamID == invalidStreamID {
|
||||
nextStreamID = 0
|
||||
} else {
|
||||
nextStreamID = c.lastStreamID + 4
|
||||
}
|
||||
c.streamMx.Unlock()
|
||||
// Streams with stream ID equal to or greater than the stream ID carried in the GOAWAY frame
|
||||
// will be rejected, see section 5.2 of RFC 9114.
|
||||
if maxStreamID != invalidStreamID && nextStreamID >= maxStreamID {
|
||||
return nil, errGoAway
|
||||
}
|
||||
|
||||
str, err := c.conn.OpenStreamSync(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
c.streamMx.Lock()
|
||||
// take the maximum here, as multiple OpenStreamSync calls might have returned concurrently
|
||||
if c.lastStreamID == invalidStreamID {
|
||||
c.lastStreamID = str.StreamID()
|
||||
} else {
|
||||
c.lastStreamID = max(c.lastStreamID, str.StreamID())
|
||||
}
|
||||
// check again, in case a (or another) GOAWAY frame was received
|
||||
maxStreamID = c.maxStreamID
|
||||
c.streamMx.Unlock()
|
||||
|
||||
if maxStreamID != invalidStreamID && str.StreamID() >= maxStreamID {
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled))
|
||||
str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestCanceled))
|
||||
return nil, errGoAway
|
||||
}
|
||||
|
||||
hstr := c.rawConn.TrackStream(str)
|
||||
rsp := &http.Response{}
|
||||
trace := httptrace.ContextClientTrace(ctx)
|
||||
return newRequestStream(
|
||||
newStream(hstr, c.rawConn, trace, func(r io.Reader, hf *headersFrame) error {
|
||||
hdr, err := decodeTrailers(r, hf, maxHeaderBytes, c.decoder, c.qlogger, str.StreamID())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rsp.Trailer = hdr
|
||||
return nil
|
||||
}, c.qlogger),
|
||||
requestWriter,
|
||||
reqDone,
|
||||
c.decoder,
|
||||
disableCompression,
|
||||
maxHeaderBytes,
|
||||
rsp,
|
||||
), nil
|
||||
}
|
||||
|
||||
func (c *ClientConn) handleUnidirectionalStream(str *quic.ReceiveStream) {
|
||||
c.rawConn.handleUnidirectionalStream(str, false)
|
||||
}
|
||||
|
||||
func (c *ClientConn) handleControlStream(str *quic.ReceiveStream, fp *frameParser) {
|
||||
for {
|
||||
f, err := fp.ParseNext(c.qlogger)
|
||||
if err != nil {
|
||||
var serr *quic.StreamError
|
||||
if err == io.EOF || errors.As(err, &serr) {
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeClosedCriticalStream), "")
|
||||
return
|
||||
}
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameError), "")
|
||||
return
|
||||
}
|
||||
// GOAWAY is the only frame allowed at this point:
|
||||
// * unexpected frames are ignored by the frame parser
|
||||
// * we don't support any extension that might add support for more frames
|
||||
goaway, ok := f.(*goAwayFrame)
|
||||
if !ok {
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "")
|
||||
return
|
||||
}
|
||||
if goaway.StreamID%4 != 0 { // client-initiated, bidirectional streams
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeIDError), "")
|
||||
return
|
||||
}
|
||||
c.streamMx.Lock()
|
||||
// the server is not allowed to increase the Stream ID in subsequent GOAWAY frames
|
||||
if c.maxStreamID != invalidStreamID && goaway.StreamID > c.maxStreamID {
|
||||
c.streamMx.Unlock()
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeIDError), "")
|
||||
return
|
||||
}
|
||||
c.maxStreamID = goaway.StreamID
|
||||
c.streamMx.Unlock()
|
||||
|
||||
hasActiveStreams := c.rawConn.hasActiveStreams()
|
||||
// immediately close the connection if there are currently no active requests
|
||||
if !hasActiveStreams {
|
||||
c.CloseWithError(quic.ApplicationErrorCode(ErrCodeNoError), "")
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *ClientConn) onStreamsEmpty() {
|
||||
c.streamMx.Lock()
|
||||
defer c.streamMx.Unlock()
|
||||
|
||||
// The server is performing a graceful shutdown.
|
||||
if c.maxStreamID != invalidStreamID {
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeNoError), "")
|
||||
}
|
||||
}
|
||||
|
||||
// RoundTrip executes a request and returns a response
|
||||
func (c *ClientConn) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
rsp, err := c.roundTrip(req)
|
||||
if err != nil && req.Context().Err() != nil {
|
||||
// if the context was canceled, return the context cancellation error
|
||||
err = req.Context().Err()
|
||||
}
|
||||
return rsp, err
|
||||
}
|
||||
|
||||
func (c *ClientConn) roundTrip(req *http.Request) (*http.Response, error) {
|
||||
// Immediately send out this request, if this is a 0-RTT request.
|
||||
switch req.Method {
|
||||
case MethodGet0RTT:
|
||||
// don't modify the original request
|
||||
reqCopy := *req
|
||||
req = &reqCopy
|
||||
req.Method = http.MethodGet
|
||||
case MethodHead0RTT:
|
||||
// don't modify the original request
|
||||
reqCopy := *req
|
||||
req = &reqCopy
|
||||
req.Method = http.MethodHead
|
||||
default:
|
||||
// wait for the handshake to complete
|
||||
select {
|
||||
case <-c.conn.HandshakeComplete():
|
||||
case <-req.Context().Done():
|
||||
return nil, req.Context().Err()
|
||||
}
|
||||
}
|
||||
|
||||
// It is only possible to send an Extended CONNECT request once the SETTINGS were received.
|
||||
// See section 3 of RFC 8441.
|
||||
if isExtendedConnectRequest(req) {
|
||||
connCtx := c.conn.Context()
|
||||
// wait for the server's SETTINGS frame to arrive
|
||||
select {
|
||||
case <-c.rawConn.ReceivedSettings():
|
||||
case <-connCtx.Done():
|
||||
return nil, context.Cause(connCtx)
|
||||
}
|
||||
if !c.rawConn.Settings().EnableExtendedConnect {
|
||||
return nil, errors.New("http3: server didn't enable Extended CONNECT")
|
||||
}
|
||||
}
|
||||
|
||||
reqDone := make(chan struct{})
|
||||
str, err := c.openRequestStream(
|
||||
req.Context(),
|
||||
c.requestWriter,
|
||||
reqDone,
|
||||
c.disableCompression,
|
||||
c.maxResponseHeaderBytes,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, &errConnUnusable{e: err}
|
||||
}
|
||||
|
||||
// Request Cancellation:
|
||||
// This go routine keeps running even after RoundTripOpt() returns.
|
||||
// It is shut down when the application is done processing the body.
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
defer close(done)
|
||||
select {
|
||||
case <-req.Context().Done():
|
||||
str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestCanceled))
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeRequestCanceled))
|
||||
case <-reqDone:
|
||||
}
|
||||
}()
|
||||
|
||||
rsp, err := c.doRequest(req, str)
|
||||
if err != nil { // if any error occurred
|
||||
close(reqDone)
|
||||
<-done
|
||||
return nil, maybeReplaceError(err)
|
||||
}
|
||||
return rsp, maybeReplaceError(err)
|
||||
}
|
||||
|
||||
// ReceivedSettings returns a channel that is closed once the server's HTTP/3 settings were received.
|
||||
// Settings can be obtained from the Settings method after the channel was closed.
|
||||
func (c *ClientConn) ReceivedSettings() <-chan struct{} {
|
||||
return c.rawConn.ReceivedSettings()
|
||||
}
|
||||
|
||||
// Settings returns the HTTP/3 settings for this connection.
|
||||
// It is only valid to call this function after the channel returned by ReceivedSettings was closed.
|
||||
func (c *ClientConn) Settings() *Settings {
|
||||
return c.rawConn.Settings()
|
||||
}
|
||||
|
||||
// CloseWithError closes the connection with the given error code and message.
|
||||
// It is invalid to call this function after the connection was closed.
|
||||
func (c *ClientConn) CloseWithError(code quic.ApplicationErrorCode, msg string) error {
|
||||
return c.conn.CloseWithError(code, msg)
|
||||
}
|
||||
|
||||
// Context returns a context that is cancelled when the connection is closed.
|
||||
func (c *ClientConn) Context() context.Context {
|
||||
return c.conn.Context()
|
||||
}
|
||||
|
||||
// cancelingReader reads from the io.Reader.
|
||||
// It cancels writing on the stream if any error other than io.EOF occurs.
|
||||
type cancelingReader struct {
|
||||
r io.Reader
|
||||
str *RequestStream
|
||||
}
|
||||
|
||||
func (r *cancelingReader) Read(b []byte) (int, error) {
|
||||
n, err := r.r.Read(b)
|
||||
if err != nil && err != io.EOF {
|
||||
r.str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestCanceled))
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (c *ClientConn) sendRequestBody(str *RequestStream, body io.ReadCloser, contentLength int64) error {
|
||||
defer body.Close()
|
||||
buf := make([]byte, bodyCopyBufferSize)
|
||||
sr := &cancelingReader{str: str, r: body}
|
||||
if contentLength == -1 {
|
||||
_, err := io.CopyBuffer(str, sr, buf)
|
||||
return err
|
||||
}
|
||||
|
||||
// make sure we don't send more bytes than the content length
|
||||
n, err := io.CopyBuffer(str, io.LimitReader(sr, contentLength), buf)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var extra int64
|
||||
extra, err = io.CopyBuffer(io.Discard, sr, buf)
|
||||
n += extra
|
||||
if n > contentLength {
|
||||
str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestCanceled))
|
||||
return fmt.Errorf("http: ContentLength=%d with Body length %d", contentLength, n)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *ClientConn) doRequest(req *http.Request, str *RequestStream) (*http.Response, error) {
|
||||
trace := httptrace.ContextClientTrace(req.Context())
|
||||
var sendingReqFailed bool
|
||||
if err := str.sendRequestHeader(req); err != nil {
|
||||
traceWroteRequest(trace, err)
|
||||
if c.logger != nil {
|
||||
c.logger.Debug("error writing request", "error", err)
|
||||
}
|
||||
sendingReqFailed = true
|
||||
}
|
||||
if !sendingReqFailed {
|
||||
if req.Body == nil {
|
||||
traceWroteRequest(trace, nil)
|
||||
str.Close()
|
||||
} else {
|
||||
// send the request body asynchronously
|
||||
go func() {
|
||||
defer str.Close()
|
||||
contentLength := int64(-1)
|
||||
// According to the documentation for http.Request.ContentLength,
|
||||
// a value of 0 with a non-nil Body is also treated as unknown content length.
|
||||
if req.ContentLength > 0 {
|
||||
contentLength = req.ContentLength
|
||||
}
|
||||
err := c.sendRequestBody(str, req.Body, contentLength)
|
||||
traceWroteRequest(trace, err)
|
||||
if err != nil {
|
||||
if c.logger != nil {
|
||||
c.logger.Debug("error writing request", "error", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
if len(req.Trailer) > 0 {
|
||||
if err := str.sendRequestTrailer(req); err != nil {
|
||||
if c.logger != nil {
|
||||
c.logger.Debug("error writing trailers", "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
// copy from net/http: support 1xx responses
|
||||
var num1xx int // number of informational 1xx headers received
|
||||
var res *http.Response
|
||||
for {
|
||||
var err error
|
||||
res, err = str.ReadResponse()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resCode := res.StatusCode
|
||||
is1xx := 100 <= resCode && resCode <= 199
|
||||
// treat 101 as a terminal status, see https://github.com/golang/go/issues/26161
|
||||
is1xxNonTerminal := is1xx && resCode != http.StatusSwitchingProtocols
|
||||
if is1xxNonTerminal {
|
||||
num1xx++
|
||||
if num1xx > max1xxResponses {
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeExcessiveLoad))
|
||||
str.CancelWrite(quic.StreamErrorCode(ErrCodeExcessiveLoad))
|
||||
return nil, errors.New("http3: too many 1xx informational responses")
|
||||
}
|
||||
traceGot1xxResponse(trace, resCode, textproto.MIMEHeader(res.Header))
|
||||
if resCode == http.StatusContinue {
|
||||
traceGot100Continue(trace)
|
||||
}
|
||||
continue
|
||||
}
|
||||
break
|
||||
}
|
||||
connState := c.conn.ConnectionState().TLS
|
||||
res.TLS = &connState
|
||||
res.Request = req
|
||||
return res, nil
|
||||
}
|
||||
|
||||
// RawClientConn is a low-level HTTP/3 client connection.
|
||||
// It allows the application to take control of the stream accept loops,
|
||||
// giving the application the ability to handle streams originating from the server.
|
||||
type RawClientConn struct {
|
||||
*ClientConn
|
||||
}
|
||||
|
||||
// HandleUnidirectionalStream handles an incoming unidirectional stream.
|
||||
func (c *RawClientConn) HandleUnidirectionalStream(str *quic.ReceiveStream) {
|
||||
c.rawConn.handleUnidirectionalStream(str, false)
|
||||
}
|
||||
|
||||
// HandleBidirectionalStream handles an incoming bidirectional stream.
|
||||
func (c *ClientConn) HandleBidirectionalStream(str *quic.Stream) {
|
||||
// According to RFC 9114, the server is not allowed to open bidirectional streams.
|
||||
c.rawConn.CloseWithError(
|
||||
quic.ApplicationErrorCode(ErrCodeStreamCreationError),
|
||||
fmt.Sprintf("server opened bidirectional stream %d", str.StreamID()),
|
||||
)
|
||||
}
|
||||
+318
@@ -0,0 +1,318 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/enetx/http3/qlog"
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
"github.com/quic-go/quic-go/quicvarint"
|
||||
)
|
||||
|
||||
const maxQuarterStreamID = 1<<60 - 1
|
||||
|
||||
// invalidStreamID is a stream ID that is invalid. The first valid stream ID in QUIC is 0.
|
||||
const invalidStreamID = quic.StreamID(-1)
|
||||
|
||||
// rawConn is an HTTP/3 connection.
|
||||
// It provides HTTP/3 specific functionality by wrapping a quic.Conn,
|
||||
// in particular handling of unidirectional HTTP/3 streams, SETTINGS and datagrams.
|
||||
type rawConn struct {
|
||||
conn *quic.Conn
|
||||
|
||||
logger *slog.Logger
|
||||
|
||||
enableDatagrams bool
|
||||
|
||||
streamMx sync.Mutex
|
||||
streams map[quic.StreamID]*stateTrackingStream
|
||||
|
||||
rcvdControlStr atomic.Bool
|
||||
rcvdQPACKEncoderStr atomic.Bool
|
||||
rcvdQPACKDecoderStr atomic.Bool
|
||||
controlStrHandler func(*quic.ReceiveStream, *frameParser) // is called *after* the SETTINGS frame was parsed
|
||||
|
||||
onStreamsEmpty func()
|
||||
|
||||
settings *Settings
|
||||
receivedSettings chan struct{}
|
||||
|
||||
qlogger qlogwriter.Recorder
|
||||
qloggerWG sync.WaitGroup // tracks goroutines that may produce qlog events
|
||||
}
|
||||
|
||||
func newRawConn(
|
||||
quicConn *quic.Conn,
|
||||
enableDatagrams bool,
|
||||
onStreamsEmpty func(),
|
||||
controlStrHandler func(*quic.ReceiveStream, *frameParser),
|
||||
qlogger qlogwriter.Recorder,
|
||||
logger *slog.Logger,
|
||||
) *rawConn {
|
||||
c := &rawConn{
|
||||
conn: quicConn,
|
||||
logger: logger,
|
||||
enableDatagrams: enableDatagrams,
|
||||
receivedSettings: make(chan struct{}),
|
||||
streams: make(map[quic.StreamID]*stateTrackingStream),
|
||||
qlogger: qlogger,
|
||||
onStreamsEmpty: onStreamsEmpty,
|
||||
controlStrHandler: controlStrHandler,
|
||||
}
|
||||
if qlogger != nil {
|
||||
context.AfterFunc(quicConn.Context(), c.closeQlogger)
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *rawConn) OpenUniStream() (*quic.SendStream, error) {
|
||||
return c.conn.OpenUniStream()
|
||||
}
|
||||
|
||||
// openControlStream opens the control stream and sends the SETTINGS frame.
|
||||
// It returns the control stream (needed by the server for sending GOAWAY later).
|
||||
func (c *rawConn) openControlStream(settings *settingsFrame) (*quic.SendStream, error) {
|
||||
c.qloggerWG.Add(1)
|
||||
defer c.qloggerWG.Done()
|
||||
|
||||
str, err := c.conn.OpenUniStream()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
b := make([]byte, 0, 64)
|
||||
b = quicvarint.Append(b, streamTypeControlStream)
|
||||
b = settings.Append(b)
|
||||
if c.qlogger != nil {
|
||||
sf := qlog.SettingsFrame{
|
||||
MaxFieldSectionSize: settings.MaxFieldSectionSize,
|
||||
Other: settings.Other.Iter().Collect().Map[uint64, uint64](),
|
||||
}
|
||||
if settings.Datagram {
|
||||
sf.Datagram = pointer(true)
|
||||
}
|
||||
if settings.ExtendedConnect {
|
||||
sf.ExtendedConnect = pointer(true)
|
||||
}
|
||||
c.qlogger.RecordEvent(qlog.FrameCreated{
|
||||
StreamID: str.StreamID(),
|
||||
Raw: qlog.RawInfo{Length: len(b)},
|
||||
Frame: qlog.Frame{Frame: sf},
|
||||
})
|
||||
}
|
||||
if _, err := str.Write(b); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return str, nil
|
||||
}
|
||||
|
||||
func (c *rawConn) TrackStream(str *quic.Stream) *stateTrackingStream {
|
||||
hstr := newStateTrackingStream(str, c, func(b []byte) error { return c.sendDatagram(str.StreamID(), b) })
|
||||
|
||||
c.streamMx.Lock()
|
||||
c.streams[str.StreamID()] = hstr
|
||||
c.qloggerWG.Add(1)
|
||||
c.streamMx.Unlock()
|
||||
return hstr
|
||||
}
|
||||
|
||||
func (c *rawConn) RemoteAddr() net.Addr {
|
||||
return c.conn.RemoteAddr()
|
||||
}
|
||||
|
||||
func (c *rawConn) ConnectionState() quic.ConnectionState {
|
||||
return c.conn.ConnectionState()
|
||||
}
|
||||
|
||||
func (c *rawConn) clearStream(id quic.StreamID) {
|
||||
c.streamMx.Lock()
|
||||
defer c.streamMx.Unlock()
|
||||
|
||||
if _, ok := c.streams[id]; ok {
|
||||
delete(c.streams, id)
|
||||
c.qloggerWG.Done()
|
||||
}
|
||||
if len(c.streams) == 0 {
|
||||
c.onStreamsEmpty()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *rawConn) hasActiveStreams() bool {
|
||||
c.streamMx.Lock()
|
||||
defer c.streamMx.Unlock()
|
||||
|
||||
return len(c.streams) > 0
|
||||
}
|
||||
|
||||
func (c *rawConn) CloseWithError(code quic.ApplicationErrorCode, msg string) error {
|
||||
return c.conn.CloseWithError(code, msg)
|
||||
}
|
||||
|
||||
func (c *rawConn) handleUnidirectionalStream(str *quic.ReceiveStream, isServer bool) {
|
||||
c.qloggerWG.Add(1)
|
||||
defer c.qloggerWG.Done()
|
||||
|
||||
streamType, err := quicvarint.Read(quicvarint.NewReader(str))
|
||||
if err != nil {
|
||||
if c.logger != nil {
|
||||
c.logger.Debug("reading stream type on stream failed", "stream ID", str.StreamID(), "error", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
// We're only interested in the control stream here.
|
||||
switch streamType {
|
||||
case streamTypeControlStream:
|
||||
case streamTypeQPACKEncoderStream:
|
||||
if isFirst := c.rcvdQPACKEncoderStr.CompareAndSwap(false, true); !isFirst {
|
||||
c.CloseWithError(quic.ApplicationErrorCode(ErrCodeStreamCreationError), "duplicate QPACK encoder stream")
|
||||
}
|
||||
// Our QPACK implementation doesn't use the dynamic table yet.
|
||||
return
|
||||
case streamTypeQPACKDecoderStream:
|
||||
if isFirst := c.rcvdQPACKDecoderStr.CompareAndSwap(false, true); !isFirst {
|
||||
c.CloseWithError(quic.ApplicationErrorCode(ErrCodeStreamCreationError), "duplicate QPACK decoder stream")
|
||||
}
|
||||
// Our QPACK implementation doesn't use the dynamic table yet.
|
||||
return
|
||||
case streamTypePushStream:
|
||||
if isServer {
|
||||
// only the server can push
|
||||
c.CloseWithError(quic.ApplicationErrorCode(ErrCodeStreamCreationError), "")
|
||||
} else {
|
||||
// we never increased the Push ID, so we don't expect any push streams
|
||||
c.CloseWithError(quic.ApplicationErrorCode(ErrCodeIDError), "")
|
||||
}
|
||||
return
|
||||
default:
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeStreamCreationError))
|
||||
return
|
||||
}
|
||||
// Only a single control stream is allowed.
|
||||
if isFirstControlStr := c.rcvdControlStr.CompareAndSwap(false, true); !isFirstControlStr {
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeStreamCreationError), "duplicate control stream")
|
||||
return
|
||||
}
|
||||
c.handleControlStream(str)
|
||||
}
|
||||
|
||||
func (c *rawConn) handleControlStream(str *quic.ReceiveStream) {
|
||||
fp := &frameParser{closeConn: c.conn.CloseWithError, r: str, streamID: str.StreamID()}
|
||||
f, err := fp.ParseNext(c.qlogger)
|
||||
if err != nil {
|
||||
var serr *quic.StreamError
|
||||
if err == io.EOF || errors.As(err, &serr) {
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeClosedCriticalStream), "")
|
||||
return
|
||||
}
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameError), "")
|
||||
return
|
||||
}
|
||||
sf, ok := f.(*settingsFrame)
|
||||
if !ok {
|
||||
c.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeMissingSettings), "")
|
||||
return
|
||||
}
|
||||
c.settings = &Settings{
|
||||
EnableDatagrams: sf.Datagram,
|
||||
EnableExtendedConnect: sf.ExtendedConnect,
|
||||
Other: sf.Other,
|
||||
}
|
||||
close(c.receivedSettings)
|
||||
if sf.Datagram {
|
||||
// If datagram support was enabled on our side as well as on the server side,
|
||||
// we can expect it to have been negotiated both on the transport and on the HTTP/3 layer.
|
||||
// Note: ConnectionState() will block until the handshake is complete (relevant when using 0-RTT).
|
||||
if c.enableDatagrams && !c.ConnectionState().SupportsDatagrams.Remote {
|
||||
c.CloseWithError(quic.ApplicationErrorCode(ErrCodeSettingsError), "missing QUIC Datagram support")
|
||||
return
|
||||
}
|
||||
c.qloggerWG.Go(func() {
|
||||
if err := c.receiveDatagrams(); err != nil {
|
||||
if c.logger != nil {
|
||||
c.logger.Debug("receiving datagrams failed", "error", err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
if c.controlStrHandler != nil {
|
||||
c.controlStrHandler(str, fp)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *rawConn) sendDatagram(streamID quic.StreamID, b []byte) error {
|
||||
// TODO: this creates a lot of garbage and an additional copy
|
||||
data := make([]byte, 0, len(b)+8)
|
||||
quarterStreamID := uint64(streamID / 4)
|
||||
data = quicvarint.Append(data, uint64(streamID/4))
|
||||
data = append(data, b...)
|
||||
if c.qlogger != nil {
|
||||
c.qlogger.RecordEvent(qlog.DatagramCreated{
|
||||
QuarterStreamID: quarterStreamID,
|
||||
Raw: qlog.RawInfo{
|
||||
Length: len(data),
|
||||
PayloadLength: len(b),
|
||||
},
|
||||
})
|
||||
}
|
||||
return c.conn.SendDatagram(data)
|
||||
}
|
||||
|
||||
func (c *rawConn) receiveDatagrams() error {
|
||||
for {
|
||||
b, err := c.conn.ReceiveDatagram(context.Background())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
quarterStreamID, n, err := quicvarint.Parse(b)
|
||||
if err != nil {
|
||||
c.CloseWithError(quic.ApplicationErrorCode(ErrCodeDatagramError), "")
|
||||
return fmt.Errorf("could not read quarter stream id: %w", err)
|
||||
}
|
||||
if c.qlogger != nil {
|
||||
c.qlogger.RecordEvent(qlog.DatagramParsed{
|
||||
QuarterStreamID: quarterStreamID,
|
||||
Raw: qlog.RawInfo{
|
||||
Length: len(b),
|
||||
PayloadLength: len(b) - n,
|
||||
},
|
||||
})
|
||||
}
|
||||
if quarterStreamID > maxQuarterStreamID {
|
||||
c.CloseWithError(quic.ApplicationErrorCode(ErrCodeDatagramError), "")
|
||||
return fmt.Errorf("invalid quarter stream id: %w", err)
|
||||
}
|
||||
streamID := quic.StreamID(4 * quarterStreamID)
|
||||
c.streamMx.Lock()
|
||||
dg, ok := c.streams[streamID]
|
||||
c.streamMx.Unlock()
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
dg.enqueueDatagram(b[n:])
|
||||
}
|
||||
}
|
||||
|
||||
// ReceivedSettings returns a channel that is closed once the peer's SETTINGS frame was received.
|
||||
// Settings can be optained from the Settings method after the channel was closed.
|
||||
func (c *rawConn) ReceivedSettings() <-chan struct{} { return c.receivedSettings }
|
||||
|
||||
// Settings returns the settings received on this connection.
|
||||
// It is only valid to call this function after the channel returned by ReceivedSettings was closed.
|
||||
func (c *rawConn) Settings() *Settings { return c.settings }
|
||||
|
||||
// closeQlogger waits for all goroutines that may produce qlog events to finish,
|
||||
// then closes the qlogger.
|
||||
func (c *rawConn) closeQlogger() {
|
||||
if c.qlogger == nil {
|
||||
return
|
||||
}
|
||||
c.qloggerWG.Wait()
|
||||
c.qlogger.Close()
|
||||
}
|
||||
+63
@@ -0,0 +1,63 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
)
|
||||
|
||||
// Error is returned from the round tripper (for HTTP clients)
|
||||
// and inside the HTTP handler (for HTTP servers) if an HTTP/3 error occurs.
|
||||
// See section 8 of RFC 9114.
|
||||
type Error struct {
|
||||
Remote bool
|
||||
ErrorCode ErrCode
|
||||
ErrorMessage string
|
||||
}
|
||||
|
||||
var _ error = &Error{}
|
||||
|
||||
func (e *Error) Error() string {
|
||||
s := e.ErrorCode.string()
|
||||
if s == "" {
|
||||
s = fmt.Sprintf("H3 error (%#x)", uint64(e.ErrorCode))
|
||||
}
|
||||
// Usually errors are remote. Only make it explicit for local errors.
|
||||
if !e.Remote {
|
||||
s += " (local)"
|
||||
}
|
||||
if e.ErrorMessage != "" {
|
||||
s += ": " + e.ErrorMessage
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func (e *Error) Is(target error) bool {
|
||||
t, ok := target.(*Error)
|
||||
return ok && e.ErrorCode == t.ErrorCode && e.Remote == t.Remote
|
||||
}
|
||||
|
||||
func maybeReplaceError(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
var (
|
||||
e Error
|
||||
strErr *quic.StreamError
|
||||
appErr *quic.ApplicationError
|
||||
)
|
||||
switch {
|
||||
default:
|
||||
return err
|
||||
case errors.As(err, &strErr):
|
||||
e.Remote = strErr.Remote
|
||||
e.ErrorCode = ErrCode(strErr.ErrorCode)
|
||||
case errors.As(err, &appErr):
|
||||
e.Remote = appErr.Remote
|
||||
e.ErrorCode = ErrCode(appErr.ErrorCode)
|
||||
e.ErrorMessage = appErr.ErrorMessage
|
||||
}
|
||||
return &e
|
||||
}
|
||||
+84
@@ -0,0 +1,84 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
)
|
||||
|
||||
type ErrCode quic.ApplicationErrorCode
|
||||
|
||||
const (
|
||||
ErrCodeNoError ErrCode = 0x100
|
||||
ErrCodeGeneralProtocolError ErrCode = 0x101
|
||||
ErrCodeInternalError ErrCode = 0x102
|
||||
ErrCodeStreamCreationError ErrCode = 0x103
|
||||
ErrCodeClosedCriticalStream ErrCode = 0x104
|
||||
ErrCodeFrameUnexpected ErrCode = 0x105
|
||||
ErrCodeFrameError ErrCode = 0x106
|
||||
ErrCodeExcessiveLoad ErrCode = 0x107
|
||||
ErrCodeIDError ErrCode = 0x108
|
||||
ErrCodeSettingsError ErrCode = 0x109
|
||||
ErrCodeMissingSettings ErrCode = 0x10a
|
||||
ErrCodeRequestRejected ErrCode = 0x10b
|
||||
ErrCodeRequestCanceled ErrCode = 0x10c
|
||||
ErrCodeRequestIncomplete ErrCode = 0x10d
|
||||
ErrCodeMessageError ErrCode = 0x10e
|
||||
ErrCodeConnectError ErrCode = 0x10f
|
||||
ErrCodeVersionFallback ErrCode = 0x110
|
||||
ErrCodeDatagramError ErrCode = 0x33
|
||||
ErrCodeQPACKDecompressionFailed ErrCode = 0x200
|
||||
)
|
||||
|
||||
func (e ErrCode) String() string {
|
||||
s := e.string()
|
||||
if s != "" {
|
||||
return s
|
||||
}
|
||||
return fmt.Sprintf("unknown error code: %#x", uint16(e))
|
||||
}
|
||||
|
||||
func (e ErrCode) string() string {
|
||||
switch e {
|
||||
case ErrCodeNoError:
|
||||
return "H3_NO_ERROR"
|
||||
case ErrCodeGeneralProtocolError:
|
||||
return "H3_GENERAL_PROTOCOL_ERROR"
|
||||
case ErrCodeInternalError:
|
||||
return "H3_INTERNAL_ERROR"
|
||||
case ErrCodeStreamCreationError:
|
||||
return "H3_STREAM_CREATION_ERROR"
|
||||
case ErrCodeClosedCriticalStream:
|
||||
return "H3_CLOSED_CRITICAL_STREAM"
|
||||
case ErrCodeFrameUnexpected:
|
||||
return "H3_FRAME_UNEXPECTED"
|
||||
case ErrCodeFrameError:
|
||||
return "H3_FRAME_ERROR"
|
||||
case ErrCodeExcessiveLoad:
|
||||
return "H3_EXCESSIVE_LOAD"
|
||||
case ErrCodeIDError:
|
||||
return "H3_ID_ERROR"
|
||||
case ErrCodeSettingsError:
|
||||
return "H3_SETTINGS_ERROR"
|
||||
case ErrCodeMissingSettings:
|
||||
return "H3_MISSING_SETTINGS"
|
||||
case ErrCodeRequestRejected:
|
||||
return "H3_REQUEST_REJECTED"
|
||||
case ErrCodeRequestCanceled:
|
||||
return "H3_REQUEST_CANCELLED"
|
||||
case ErrCodeRequestIncomplete:
|
||||
return "H3_INCOMPLETE_REQUEST"
|
||||
case ErrCodeMessageError:
|
||||
return "H3_MESSAGE_ERROR"
|
||||
case ErrCodeConnectError:
|
||||
return "H3_CONNECT_ERROR"
|
||||
case ErrCodeVersionFallback:
|
||||
return "H3_VERSION_FALLBACK"
|
||||
case ErrCodeDatagramError:
|
||||
return "H3_DATAGRAM_ERROR"
|
||||
case ErrCodeQPACKDecompressionFailed:
|
||||
return "QPACK_DECOMPRESSION_FAILED"
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
+339
@@ -0,0 +1,339 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/enetx/g"
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/http3/qlog"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
"github.com/quic-go/quic-go/quicvarint"
|
||||
)
|
||||
|
||||
// FrameType is the frame type of a HTTP/3 frame
|
||||
type FrameType uint64
|
||||
|
||||
type frame any
|
||||
|
||||
// The maximum length of an encoded HTTP/3 frame header is 16:
|
||||
// The frame has a type and length field, both QUIC varints (maximum 8 bytes in length)
|
||||
const frameHeaderLen = 16
|
||||
|
||||
type countingByteReader struct {
|
||||
quicvarint.Reader
|
||||
NumRead int
|
||||
}
|
||||
|
||||
func (r *countingByteReader) ReadByte() (byte, error) {
|
||||
b, err := r.Reader.ReadByte()
|
||||
if err == nil {
|
||||
r.NumRead++
|
||||
}
|
||||
return b, err
|
||||
}
|
||||
|
||||
func (r *countingByteReader) Read(b []byte) (int, error) {
|
||||
n, err := r.Reader.Read(b)
|
||||
r.NumRead += n
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (r *countingByteReader) Reset() {
|
||||
r.NumRead = 0
|
||||
}
|
||||
|
||||
type frameParser struct {
|
||||
r io.Reader
|
||||
streamID quic.StreamID
|
||||
closeConn func(quic.ApplicationErrorCode, string) error
|
||||
}
|
||||
|
||||
func (p *frameParser) ParseNext(qlogger qlogwriter.Recorder) (frame, error) {
|
||||
r := &countingByteReader{Reader: quicvarint.NewReader(p.r)}
|
||||
for {
|
||||
t, err := quicvarint.Read(r)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
l, err := quicvarint.Read(r)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
switch t {
|
||||
case 0x0: // DATA
|
||||
if qlogger != nil {
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: p.streamID,
|
||||
Raw: qlog.RawInfo{
|
||||
Length: int(l) + r.NumRead,
|
||||
PayloadLength: int(l),
|
||||
},
|
||||
Frame: qlog.Frame{Frame: qlog.DataFrame{}},
|
||||
})
|
||||
}
|
||||
return &dataFrame{Length: l}, nil
|
||||
case 0x1: // HEADERS
|
||||
return &headersFrame{
|
||||
Length: l,
|
||||
headerLen: r.NumRead,
|
||||
}, nil
|
||||
case 0x4: // SETTINGS
|
||||
return parseSettingsFrame(r, l, p.streamID, qlogger)
|
||||
case 0x3: // unsupported: CANCEL_PUSH
|
||||
if qlogger != nil {
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: p.streamID,
|
||||
Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)},
|
||||
Frame: qlog.Frame{Frame: qlog.CancelPushFrame{}},
|
||||
})
|
||||
}
|
||||
case 0x5: // unsupported: PUSH_PROMISE
|
||||
if qlogger != nil {
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: p.streamID,
|
||||
Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)},
|
||||
Frame: qlog.Frame{Frame: qlog.PushPromiseFrame{}},
|
||||
})
|
||||
}
|
||||
case 0x7: // GOAWAY
|
||||
return parseGoAwayFrame(r, l, p.streamID, qlogger)
|
||||
case 0xd: // unsupported: MAX_PUSH_ID
|
||||
if qlogger != nil {
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: p.streamID,
|
||||
Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)},
|
||||
Frame: qlog.Frame{Frame: qlog.MaxPushIDFrame{}},
|
||||
})
|
||||
}
|
||||
case 0x2, 0x6, 0x8, 0x9: // reserved frame types
|
||||
if qlogger != nil {
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: p.streamID,
|
||||
Raw: qlog.RawInfo{Length: r.NumRead + int(l), PayloadLength: int(l)},
|
||||
Frame: qlog.Frame{Frame: qlog.ReservedFrame{Type: t}},
|
||||
})
|
||||
}
|
||||
p.closeConn(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "")
|
||||
return nil, fmt.Errorf("http3: reserved frame type: %d", t)
|
||||
default:
|
||||
// unknown frame types
|
||||
if qlogger != nil {
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: p.streamID,
|
||||
Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)},
|
||||
Frame: qlog.Frame{Frame: qlog.UnknownFrame{Type: t}},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// skip over the payload
|
||||
if _, err := io.CopyN(io.Discard, r, int64(l)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.Reset()
|
||||
}
|
||||
}
|
||||
|
||||
type dataFrame struct {
|
||||
Length uint64
|
||||
}
|
||||
|
||||
func (f *dataFrame) Append(b []byte) []byte {
|
||||
b = quicvarint.Append(b, 0x0)
|
||||
return quicvarint.Append(b, f.Length)
|
||||
}
|
||||
|
||||
type headersFrame struct {
|
||||
Length uint64
|
||||
headerLen int // number of bytes read for type and length field
|
||||
}
|
||||
|
||||
func (f *headersFrame) Append(b []byte) []byte {
|
||||
b = quicvarint.Append(b, 0x1)
|
||||
return quicvarint.Append(b, f.Length)
|
||||
}
|
||||
|
||||
const (
|
||||
// SETTINGS_MAX_FIELD_SECTION_SIZE
|
||||
settingMaxFieldSectionSize = 0x6
|
||||
// Extended CONNECT, RFC 9220
|
||||
settingExtendedConnect = 0x8
|
||||
// HTTP Datagrams, RFC 9297
|
||||
settingDatagram = 0x33
|
||||
)
|
||||
|
||||
type settingsFrame struct {
|
||||
MaxFieldSectionSize int64 // SETTINGS_MAX_FIELD_SECTION_SIZE, -1 if not set
|
||||
|
||||
Datagram bool // HTTP Datagrams, RFC 9297
|
||||
ExtendedConnect bool // Extended CONNECT, RFC 9220
|
||||
Other g.MapOrd[uint64, uint64] // all settings that we don't explicitly recognize
|
||||
}
|
||||
|
||||
func pointer[T any](v T) *T {
|
||||
return &v
|
||||
}
|
||||
|
||||
func parseSettingsFrame(
|
||||
r *countingByteReader,
|
||||
l uint64,
|
||||
streamID quic.StreamID,
|
||||
qlogger qlogwriter.Recorder,
|
||||
) (*settingsFrame, error) {
|
||||
if l > 8*(1<<10) {
|
||||
return nil, fmt.Errorf("unexpected size for SETTINGS frame: %d", l)
|
||||
}
|
||||
buf := make([]byte, l)
|
||||
if _, err := io.ReadFull(r, buf); err != nil {
|
||||
if err == io.ErrUnexpectedEOF {
|
||||
return nil, io.EOF
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
frame := &settingsFrame{MaxFieldSectionSize: -1}
|
||||
b := bytes.NewReader(buf)
|
||||
settingsFrame := qlog.SettingsFrame{MaxFieldSectionSize: -1}
|
||||
var readMaxFieldSectionSize, readDatagram, readExtendedConnect bool
|
||||
for b.Len() > 0 {
|
||||
id, err := quicvarint.Read(b)
|
||||
if err != nil { // should not happen. We allocated the whole frame already.
|
||||
return nil, err
|
||||
}
|
||||
val, err := quicvarint.Read(b)
|
||||
if err != nil { // should not happen. We allocated the whole frame already.
|
||||
return nil, err
|
||||
}
|
||||
|
||||
switch id {
|
||||
case settingMaxFieldSectionSize:
|
||||
if readMaxFieldSectionSize {
|
||||
return nil, fmt.Errorf("duplicate setting: %d", id)
|
||||
}
|
||||
readMaxFieldSectionSize = true
|
||||
frame.MaxFieldSectionSize = int64(val)
|
||||
settingsFrame.MaxFieldSectionSize = int64(val)
|
||||
case settingExtendedConnect:
|
||||
if readExtendedConnect {
|
||||
return nil, fmt.Errorf("duplicate setting: %d", id)
|
||||
}
|
||||
readExtendedConnect = true
|
||||
if val != 0 && val != 1 {
|
||||
return nil, fmt.Errorf("invalid value for SETTINGS_ENABLE_CONNECT_PROTOCOL: %d", val)
|
||||
}
|
||||
frame.ExtendedConnect = val == 1
|
||||
if qlogger != nil {
|
||||
settingsFrame.ExtendedConnect = pointer(frame.ExtendedConnect)
|
||||
}
|
||||
case settingDatagram:
|
||||
if readDatagram {
|
||||
return nil, fmt.Errorf("duplicate setting: %d", id)
|
||||
}
|
||||
readDatagram = true
|
||||
if val != 0 && val != 1 {
|
||||
return nil, fmt.Errorf("invalid value for SETTINGS_H3_DATAGRAM: %d", val)
|
||||
}
|
||||
frame.Datagram = val == 1
|
||||
if qlogger != nil {
|
||||
settingsFrame.Datagram = pointer(frame.Datagram)
|
||||
}
|
||||
default:
|
||||
if frame.Other.Contains(id) {
|
||||
return nil, fmt.Errorf("duplicate setting: %d", id)
|
||||
}
|
||||
if frame.Other == nil {
|
||||
frame.Other = g.NewMapOrd[uint64, uint64]()
|
||||
}
|
||||
frame.Other.Insert(id, val)
|
||||
}
|
||||
}
|
||||
if qlogger != nil {
|
||||
settingsFrame.Other = frame.Other.Iter().Collect().Map[uint64, uint64]()
|
||||
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: streamID,
|
||||
Raw: qlog.RawInfo{
|
||||
Length: r.NumRead,
|
||||
PayloadLength: int(l),
|
||||
},
|
||||
Frame: qlog.Frame{Frame: settingsFrame},
|
||||
})
|
||||
}
|
||||
return frame, nil
|
||||
}
|
||||
|
||||
func (f *settingsFrame) Append(b []byte) []byte {
|
||||
b = quicvarint.Append(b, 0x4)
|
||||
var l int
|
||||
if f.MaxFieldSectionSize > 0 { // enetx
|
||||
// if f.MaxFieldSectionSize >= 0 {
|
||||
l += quicvarint.Len(settingMaxFieldSectionSize) + quicvarint.Len(uint64(f.MaxFieldSectionSize))
|
||||
}
|
||||
for id, val := range f.Other.Iter() {
|
||||
l += quicvarint.Len(id) + quicvarint.Len(val)
|
||||
}
|
||||
if f.Datagram {
|
||||
l += quicvarint.Len(settingDatagram) + quicvarint.Len(1)
|
||||
}
|
||||
if f.ExtendedConnect {
|
||||
l += quicvarint.Len(settingExtendedConnect) + quicvarint.Len(1)
|
||||
}
|
||||
b = quicvarint.Append(b, uint64(l))
|
||||
if f.MaxFieldSectionSize > 0 {
|
||||
// if f.MaxFieldSectionSize >= 0 { // enetx
|
||||
b = quicvarint.Append(b, settingMaxFieldSectionSize)
|
||||
b = quicvarint.Append(b, uint64(f.MaxFieldSectionSize))
|
||||
}
|
||||
if f.Datagram {
|
||||
b = quicvarint.Append(b, settingDatagram)
|
||||
b = quicvarint.Append(b, 1)
|
||||
}
|
||||
if f.ExtendedConnect {
|
||||
b = quicvarint.Append(b, settingExtendedConnect)
|
||||
b = quicvarint.Append(b, 1)
|
||||
}
|
||||
for id, val := range f.Other.Iter() {
|
||||
b = quicvarint.Append(b, id)
|
||||
b = quicvarint.Append(b, val)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
type goAwayFrame struct {
|
||||
StreamID quic.StreamID
|
||||
}
|
||||
|
||||
func parseGoAwayFrame(
|
||||
r *countingByteReader,
|
||||
l uint64,
|
||||
streamID quic.StreamID,
|
||||
qlogger qlogwriter.Recorder,
|
||||
) (*goAwayFrame, error) {
|
||||
frame := &goAwayFrame{}
|
||||
startLen := r.NumRead
|
||||
id, err := quicvarint.Read(r)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if r.NumRead-startLen != int(l) {
|
||||
return nil, errors.New("GOAWAY frame: inconsistent length")
|
||||
}
|
||||
frame.StreamID = quic.StreamID(id)
|
||||
if qlogger != nil {
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: streamID,
|
||||
Raw: qlog.RawInfo{Length: r.NumRead, PayloadLength: int(l)},
|
||||
Frame: qlog.Frame{Frame: qlog.GoAwayFrame{StreamID: frame.StreamID}},
|
||||
})
|
||||
}
|
||||
return frame, nil
|
||||
}
|
||||
|
||||
func (f *goAwayFrame) Append(b []byte) []byte {
|
||||
b = quicvarint.Append(b, 0x7)
|
||||
b = quicvarint.Append(b, uint64(quicvarint.Len(uint64(f.StreamID))))
|
||||
return quicvarint.Append(b, uint64(f.StreamID))
|
||||
}
|
||||
+39
@@ -0,0 +1,39 @@
|
||||
package http3
|
||||
|
||||
// copied from net/transport.go
|
||||
|
||||
// gzipReader wraps a response body so it can lazily
|
||||
// call gzip.NewReader on the first call to Read
|
||||
import (
|
||||
"compress/gzip"
|
||||
"io"
|
||||
)
|
||||
|
||||
// call gzip.NewReader on the first call to Read
|
||||
type gzipReader struct {
|
||||
body io.ReadCloser // underlying Response.Body
|
||||
zr *gzip.Reader // lazily-initialized gzip reader
|
||||
zerr error // sticky error
|
||||
}
|
||||
|
||||
func newGzipReader(body io.ReadCloser) io.ReadCloser {
|
||||
return &gzipReader{body: body}
|
||||
}
|
||||
|
||||
func (gz *gzipReader) Read(p []byte) (n int, err error) {
|
||||
if gz.zerr != nil {
|
||||
return 0, gz.zerr
|
||||
}
|
||||
if gz.zr == nil {
|
||||
gz.zr, err = gzip.NewReader(gz.body)
|
||||
if err != nil {
|
||||
gz.zerr = err
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
return gz.zr.Read(p)
|
||||
}
|
||||
|
||||
func (gz *gzipReader) Close() error {
|
||||
return gz.body.Close()
|
||||
}
|
||||
+429
@@ -0,0 +1,429 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"github.com/enetx/http"
|
||||
"net/textproto"
|
||||
"net/url"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/net/http/httpguts"
|
||||
|
||||
"github.com/quic-go/qpack"
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/enetx/http3/qlog"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
)
|
||||
|
||||
type qpackError struct{ err error }
|
||||
|
||||
func (e *qpackError) Error() string { return fmt.Sprintf("qpack: %v", e.err) }
|
||||
func (e *qpackError) Unwrap() error { return e.err }
|
||||
|
||||
var errHeaderTooLarge = errors.New("http3: headers too large")
|
||||
|
||||
type header struct {
|
||||
// Pseudo header fields defined in RFC 9114
|
||||
Path string
|
||||
Method string
|
||||
Authority string
|
||||
Scheme string
|
||||
Status string
|
||||
// for Extended connect
|
||||
Protocol string
|
||||
// parsed and deduplicated. -1 if no Content-Length header is sent
|
||||
ContentLength int64
|
||||
// all non-pseudo headers
|
||||
Headers http.Header
|
||||
}
|
||||
|
||||
// connection-specific header fields must not be sent on HTTP/3
|
||||
var invalidHeaderFields = [...]string{
|
||||
"connection",
|
||||
"keep-alive",
|
||||
"proxy-connection",
|
||||
"transfer-encoding",
|
||||
"upgrade",
|
||||
}
|
||||
|
||||
func parseHeaders(decodeFn qpack.DecodeFunc, isRequest bool, sizeLimit int, headerFields *[]qpack.HeaderField) (header, error) {
|
||||
hdr := header{Headers: make(http.Header)}
|
||||
var readFirstRegularHeader, readContentLength bool
|
||||
var contentLengthStr string
|
||||
for {
|
||||
h, err := decodeFn()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
return header{}, &qpackError{err}
|
||||
}
|
||||
if headerFields != nil {
|
||||
*headerFields = append(*headerFields, h)
|
||||
}
|
||||
// RFC 9114, section 4.2.2:
|
||||
// The size of a field list is calculated based on the uncompressed size of fields,
|
||||
// including the length of the name and value in bytes plus an overhead of 32 bytes for each field.
|
||||
sizeLimit -= len(h.Name) + len(h.Value) + 32
|
||||
if sizeLimit < 0 {
|
||||
return header{}, errHeaderTooLarge
|
||||
}
|
||||
if err := validateHeaderFieldNameAndValue(h); err != nil {
|
||||
return header{}, err
|
||||
}
|
||||
if h.IsPseudo() {
|
||||
if readFirstRegularHeader {
|
||||
// all pseudo headers must appear before regular header fields, see section 4.3 of RFC 9114
|
||||
return header{}, fmt.Errorf("received pseudo header %s after a regular header field", h.Name)
|
||||
}
|
||||
var isResponsePseudoHeader bool // pseudo headers are either valid for requests or for responses
|
||||
var isDuplicatePseudoHeader bool // pseudo headers are allowed to appear exactly once
|
||||
switch h.Name {
|
||||
case ":path":
|
||||
isDuplicatePseudoHeader = hdr.Path != ""
|
||||
hdr.Path = h.Value
|
||||
case ":method":
|
||||
isDuplicatePseudoHeader = hdr.Method != ""
|
||||
hdr.Method = h.Value
|
||||
case ":authority":
|
||||
isDuplicatePseudoHeader = hdr.Authority != ""
|
||||
hdr.Authority = h.Value
|
||||
case ":protocol": // RFC 9220
|
||||
isDuplicatePseudoHeader = hdr.Protocol != ""
|
||||
hdr.Protocol = h.Value
|
||||
case ":scheme":
|
||||
isDuplicatePseudoHeader = hdr.Scheme != ""
|
||||
hdr.Scheme = h.Value
|
||||
case ":status":
|
||||
isDuplicatePseudoHeader = hdr.Status != ""
|
||||
hdr.Status = h.Value
|
||||
isResponsePseudoHeader = true
|
||||
default:
|
||||
return header{}, fmt.Errorf("unknown pseudo header: %s", h.Name)
|
||||
}
|
||||
if isDuplicatePseudoHeader {
|
||||
return header{}, fmt.Errorf("duplicate pseudo header: %s", h.Name)
|
||||
}
|
||||
if isRequest && isResponsePseudoHeader {
|
||||
return header{}, fmt.Errorf("invalid request pseudo header: %s", h.Name)
|
||||
}
|
||||
if !isRequest && !isResponsePseudoHeader {
|
||||
return header{}, fmt.Errorf("invalid response pseudo header: %s", h.Name)
|
||||
}
|
||||
} else {
|
||||
if err := validateRegularHeaderField(h); err != nil {
|
||||
return header{}, err
|
||||
}
|
||||
readFirstRegularHeader = true
|
||||
switch h.Name {
|
||||
case "content-length":
|
||||
// Ignore duplicate Content-Length headers.
|
||||
// Fail if the duplicates differ.
|
||||
if !readContentLength {
|
||||
readContentLength = true
|
||||
contentLengthStr = h.Value
|
||||
} else if contentLengthStr != h.Value {
|
||||
return header{}, fmt.Errorf("contradicting content lengths (%s and %s)", contentLengthStr, h.Value)
|
||||
}
|
||||
default:
|
||||
hdr.Headers.Add(h.Name, h.Value)
|
||||
}
|
||||
}
|
||||
}
|
||||
hdr.ContentLength = -1
|
||||
if len(contentLengthStr) > 0 {
|
||||
// use ParseUint instead of ParseInt, so that parsing fails on negative values
|
||||
cl, err := strconv.ParseUint(contentLengthStr, 10, 63)
|
||||
if err != nil {
|
||||
return header{}, fmt.Errorf("invalid content length: %w", err)
|
||||
}
|
||||
hdr.Headers.Set("Content-Length", contentLengthStr)
|
||||
hdr.ContentLength = int64(cl)
|
||||
}
|
||||
return hdr, nil
|
||||
}
|
||||
|
||||
func validateHeaderFieldNameAndValue(h qpack.HeaderField) error {
|
||||
// field names need to be lowercase, see section 4.2 of RFC 9114
|
||||
if strings.ToLower(h.Name) != h.Name {
|
||||
return fmt.Errorf("header field is not lower-case: %s", h.Name)
|
||||
}
|
||||
if !httpguts.ValidHeaderFieldValue(h.Value) {
|
||||
return fmt.Errorf("invalid header field value for %s: %q", h.Name, h.Value)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateRegularHeaderField(h qpack.HeaderField) error {
|
||||
if !httpguts.ValidHeaderFieldName(h.Name) {
|
||||
return fmt.Errorf("invalid header field name: %q", h.Name)
|
||||
}
|
||||
if slices.Contains(invalidHeaderFields[:], h.Name) {
|
||||
return fmt.Errorf("invalid header field name: %q", h.Name)
|
||||
}
|
||||
if h.Name == "te" && h.Value != "trailers" {
|
||||
return fmt.Errorf("invalid TE header field value: %q", h.Value)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateTrailerHeaderField(h qpack.HeaderField) error {
|
||||
if err := validateRegularHeaderField(h); err != nil {
|
||||
return err
|
||||
}
|
||||
if !httpguts.ValidTrailerHeader(h.Name) {
|
||||
return fmt.Errorf("invalid trailer field name: %q", h.Name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseTrailers(decodeFn qpack.DecodeFunc, sizeLimit int, headerFields *[]qpack.HeaderField) (http.Header, error) {
|
||||
h := make(http.Header)
|
||||
for {
|
||||
hf, err := decodeFn()
|
||||
if err != nil {
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
return nil, &qpackError{err}
|
||||
}
|
||||
if headerFields != nil {
|
||||
*headerFields = append(*headerFields, hf)
|
||||
}
|
||||
// RFC 9114, section 4.2.2:
|
||||
// The size of a field list is calculated based on the uncompressed size of fields,
|
||||
// including the length of the name and value in bytes plus an overhead of 32 bytes for each field.
|
||||
sizeLimit -= len(hf.Name) + len(hf.Value) + 32
|
||||
if sizeLimit < 0 {
|
||||
return nil, errHeaderTooLarge
|
||||
}
|
||||
if err := validateHeaderFieldNameAndValue(hf); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if hf.IsPseudo() {
|
||||
return nil, fmt.Errorf("http3: received pseudo header in trailer: %s", hf.Name)
|
||||
}
|
||||
if err := validateTrailerHeaderField(hf); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
h.Add(hf.Name, hf.Value)
|
||||
}
|
||||
return h, nil
|
||||
}
|
||||
|
||||
func requestFromHeaders(decodeFn qpack.DecodeFunc, sizeLimit int, headerFields *[]qpack.HeaderField) (*http.Request, error) {
|
||||
hdr, err := parseHeaders(decodeFn, true, sizeLimit, headerFields)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// concatenate cookie headers, see https://tools.ietf.org/html/rfc6265#section-5.4
|
||||
if len(hdr.Headers["Cookie"]) > 0 {
|
||||
hdr.Headers.Set("Cookie", strings.Join(hdr.Headers["Cookie"], "; "))
|
||||
}
|
||||
|
||||
isConnect := hdr.Method == http.MethodConnect
|
||||
// Extended CONNECT, see https://datatracker.ietf.org/doc/html/rfc8441#section-4
|
||||
isExtendedConnected := isConnect && hdr.Protocol != ""
|
||||
if isExtendedConnected {
|
||||
if !validExtendedConnectProtocol(hdr.Protocol) {
|
||||
return nil, fmt.Errorf("invalid :protocol: %q", hdr.Protocol)
|
||||
}
|
||||
if hdr.Scheme == "" || hdr.Path == "" || hdr.Authority == "" {
|
||||
return nil, errors.New("extended CONNECT: :scheme, :path and :authority must not be empty")
|
||||
}
|
||||
} else if isConnect {
|
||||
if hdr.Path != "" || hdr.Authority == "" { // normal CONNECT
|
||||
return nil, errors.New(":path must be empty and :authority must not be empty")
|
||||
}
|
||||
} else if len(hdr.Path) == 0 || len(hdr.Authority) == 0 || len(hdr.Method) == 0 {
|
||||
return nil, errors.New(":path, :authority and :method must not be empty")
|
||||
}
|
||||
|
||||
if !isExtendedConnected && len(hdr.Protocol) > 0 {
|
||||
return nil, errors.New(":protocol must be empty")
|
||||
}
|
||||
|
||||
var u *url.URL
|
||||
var requestURI string
|
||||
|
||||
protocol := "HTTP/3.0"
|
||||
|
||||
if isConnect {
|
||||
u = &url.URL{}
|
||||
if isExtendedConnected {
|
||||
u, err = url.ParseRequestURI(hdr.Path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
protocol = hdr.Protocol
|
||||
} else {
|
||||
u.Path = hdr.Path
|
||||
}
|
||||
requestURI = hdr.Authority
|
||||
} else {
|
||||
u, err = url.ParseRequestURI(hdr.Path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid request URI: %w", err)
|
||||
}
|
||||
requestURI = hdr.Path
|
||||
}
|
||||
u.Scheme = hdr.Scheme
|
||||
u.Host = hdr.Authority
|
||||
|
||||
req := &http.Request{
|
||||
Method: hdr.Method,
|
||||
URL: u,
|
||||
Proto: protocol,
|
||||
ProtoMajor: 3,
|
||||
ProtoMinor: 0,
|
||||
Header: hdr.Headers,
|
||||
Body: nil,
|
||||
ContentLength: hdr.ContentLength,
|
||||
Host: hdr.Authority,
|
||||
RequestURI: requestURI,
|
||||
}
|
||||
req.Trailer = extractAnnouncedTrailers(req.Header)
|
||||
return req, nil
|
||||
}
|
||||
|
||||
func validExtendedConnectProtocol(protocol string) bool {
|
||||
// RFC 9220 specifies that the semantics of the :protocol pseudo are the same as defined in RFC 8441.
|
||||
// RFC 8441, Section 4 specifies that :protocol is a single value from the HTTP Upgrade Token Registry.
|
||||
// RFC 9110, Section 16.7 specifies that HTTP Upgrade Token Registry uses token grammar.
|
||||
// Therefore, ValidHeaderFieldName is the right syntax check here, despite the misleading name.
|
||||
return httpguts.ValidHeaderFieldName(protocol)
|
||||
}
|
||||
|
||||
// updateResponseFromHeaders sets up http.Response as an HTTP/3 response,
|
||||
// using the decoded qpack header filed.
|
||||
// It is only called for the HTTP header (and not the HTTP trailer).
|
||||
// It takes an http.Response as an argument to allow the caller to set the trailer later on.
|
||||
func updateResponseFromHeaders(rsp *http.Response, decodeFn qpack.DecodeFunc, sizeLimit int, headerFields *[]qpack.HeaderField) error {
|
||||
hdr, err := parseHeaders(decodeFn, false, sizeLimit, headerFields)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if hdr.Status == "" {
|
||||
return errors.New("missing :status field")
|
||||
}
|
||||
rsp.Proto = "HTTP/3.0"
|
||||
rsp.ProtoMajor = 3
|
||||
rsp.Header = hdr.Headers
|
||||
rsp.Trailer = extractAnnouncedTrailers(rsp.Header)
|
||||
rsp.ContentLength = hdr.ContentLength
|
||||
|
||||
status, err := strconv.Atoi(hdr.Status)
|
||||
if err != nil {
|
||||
return fmt.Errorf("invalid status code: %w", err)
|
||||
}
|
||||
rsp.StatusCode = status
|
||||
rsp.Status = hdr.Status + " " + http.StatusText(status)
|
||||
return nil
|
||||
}
|
||||
|
||||
// extractAnnouncedTrailers extracts trailer keys from the "Trailer" header.
|
||||
// It returns a map with the announced keys set to nil values, and removes the "Trailer" header.
|
||||
// It handles both duplicate as well as comma-separated values for the Trailer header.
|
||||
// For example:
|
||||
//
|
||||
// Trailer: Trailer1, Trailer2
|
||||
// Trailer: Trailer3
|
||||
//
|
||||
// Will result in a map containing the keys "Trailer1", "Trailer2", "Trailer3" with nil values.
|
||||
func extractAnnouncedTrailers(header http.Header) http.Header {
|
||||
rawTrailers, ok := header["Trailer"]
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
|
||||
trailers := make(http.Header)
|
||||
for _, rawVal := range rawTrailers {
|
||||
for val := range strings.SplitSeq(rawVal, ",") {
|
||||
trailers[http.CanonicalHeaderKey(textproto.TrimString(val))] = nil
|
||||
}
|
||||
}
|
||||
delete(header, "Trailer")
|
||||
return trailers
|
||||
}
|
||||
|
||||
// writeTrailers encodes and writes HTTP trailers as a HEADERS frame.
|
||||
// It returns true if trailers were written, false if there were no trailers to write.
|
||||
func writeTrailers(wr io.Writer, trailers http.Header, streamID quic.StreamID, qlogger qlogwriter.Recorder) (bool, error) {
|
||||
var hasValues bool
|
||||
for k, vals := range trailers {
|
||||
if httpguts.ValidTrailerHeader(k) && len(vals) > 0 {
|
||||
hasValues = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !hasValues {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
enc := qpack.NewEncoder(&buf)
|
||||
var headerFields []qlog.HeaderField
|
||||
if qlogger != nil {
|
||||
headerFields = make([]qlog.HeaderField, 0, len(trailers))
|
||||
}
|
||||
|
||||
for k, vals := range trailers {
|
||||
if len(vals) == 0 {
|
||||
continue
|
||||
}
|
||||
if !httpguts.ValidTrailerHeader(k) {
|
||||
continue
|
||||
}
|
||||
lowercaseKey := strings.ToLower(k)
|
||||
for _, v := range vals {
|
||||
if err := enc.WriteField(qpack.HeaderField{Name: lowercaseKey, Value: v}); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if qlogger != nil {
|
||||
headerFields = append(headerFields, qlog.HeaderField{Name: lowercaseKey, Value: v})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
b := make([]byte, 0, frameHeaderLen+buf.Len())
|
||||
b = (&headersFrame{Length: uint64(buf.Len())}).Append(b)
|
||||
b = append(b, buf.Bytes()...)
|
||||
if qlogger != nil {
|
||||
qlogCreatedHeadersFrame(qlogger, streamID, len(b), buf.Len(), headerFields)
|
||||
}
|
||||
_, err := wr.Write(b)
|
||||
return true, err
|
||||
}
|
||||
|
||||
func decodeTrailers(r io.Reader, hf *headersFrame, maxHeaderBytes int, decoder *qpack.Decoder, qlogger qlogwriter.Recorder, streamID quic.StreamID) (http.Header, error) {
|
||||
if hf.Length > uint64(maxHeaderBytes) {
|
||||
maybeQlogInvalidHeadersFrame(qlogger, streamID, hf.Length)
|
||||
return nil, fmt.Errorf("http3: HEADERS frame too large: %d bytes (max: %d)", hf.Length, maxHeaderBytes)
|
||||
}
|
||||
|
||||
b := make([]byte, hf.Length)
|
||||
if _, err := io.ReadFull(r, b); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
decodeFn := decoder.Decode(b)
|
||||
var fields []qpack.HeaderField
|
||||
var headerFields *[]qpack.HeaderField
|
||||
if qlogger != nil {
|
||||
fields = make([]qpack.HeaderField, 0, 16)
|
||||
headerFields = &fields
|
||||
}
|
||||
trailers, err := parseTrailers(decodeFn, maxHeaderBytes, headerFields)
|
||||
if err != nil {
|
||||
maybeQlogInvalidHeadersFrame(qlogger, streamID, hf.Length)
|
||||
return nil, err
|
||||
}
|
||||
if qlogger != nil {
|
||||
qlogParsedHeadersFrame(qlogger, streamID, hf, fields)
|
||||
}
|
||||
return trailers, nil
|
||||
}
|
||||
+130
@@ -0,0 +1,130 @@
|
||||
package httpcommon
|
||||
|
||||
import (
|
||||
"github.com/enetx/http"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
const (
|
||||
// HeaderOrderKey is a magic key for ResponseWriter.Header map keys
|
||||
// that, if present, defines a header order that will be used to
|
||||
// write the headers onto wire. The order of the list defined how the headers
|
||||
// will be sorted. A defined key goes before an undefined key.
|
||||
//
|
||||
// This is the only way to specify some order, because maps don't
|
||||
// have a a stable iteration order. If no order is given, headers will
|
||||
// be sorted lexicographically.
|
||||
//
|
||||
// According to RFC-2616 it is good practice to send general-header fields
|
||||
// first, followed by request-header or response-header fields and ending
|
||||
// with entity-header fields.
|
||||
HeaderOrderKey = "Header-Order:"
|
||||
|
||||
// PHeaderOrderKey is a magic key for setting http3 pseudo header order.
|
||||
// If the header is nil it will use regular GoLang header order.
|
||||
// Valid fields are :authority, :method, :path, :scheme, :protocol
|
||||
PHeaderOrderKey = "PHeader-Order:"
|
||||
)
|
||||
|
||||
// HeaderKeyValues represents a key-value pair for headers
|
||||
type HeaderKeyValues struct {
|
||||
Key string
|
||||
Values []string
|
||||
}
|
||||
|
||||
// A HeaderSorter implements sort.Interface by sorting a []HeaderKeyValues
|
||||
// by key. It's used as a pointer, so it can fit in a sort.Interface
|
||||
// interface value without allocation.
|
||||
type HeaderSorter struct {
|
||||
kvs []HeaderKeyValues
|
||||
order map[string]int
|
||||
}
|
||||
|
||||
func (s *HeaderSorter) Len() int { return len(s.kvs) }
|
||||
func (s *HeaderSorter) Swap(i, j int) { s.kvs[i], s.kvs[j] = s.kvs[j], s.kvs[i] }
|
||||
func (s *HeaderSorter) Less(i, j int) bool {
|
||||
// If the order isn't defined, sort lexicographically.
|
||||
if s.order == nil {
|
||||
return s.kvs[i].Key < s.kvs[j].Key
|
||||
}
|
||||
|
||||
idxi, iok := s.order[strings.ToLower(s.kvs[i].Key)]
|
||||
idxj, jok := s.order[strings.ToLower(s.kvs[j].Key)]
|
||||
if !iok && !jok {
|
||||
return s.kvs[i].Key < s.kvs[j].Key
|
||||
} else if !iok && jok {
|
||||
return false
|
||||
} else if iok && !jok {
|
||||
return true
|
||||
}
|
||||
|
||||
return idxi < idxj
|
||||
}
|
||||
|
||||
var headerSorterPool = sync.Pool{
|
||||
New: func() any { return new(HeaderSorter) },
|
||||
}
|
||||
|
||||
var lock = sync.RWMutex{}
|
||||
|
||||
// SortedKeyValues returns h's keys sorted in the returned kvs
|
||||
// slice. The HeaderSorter used to sort is also returned, for possible
|
||||
// return to headerSorterPool.
|
||||
func SortedKeyValues(h http.Header, exclude map[string]bool) (kvs []HeaderKeyValues, hs *HeaderSorter) {
|
||||
hs = headerSorterPool.Get().(*HeaderSorter)
|
||||
if cap(hs.kvs) < len(h) {
|
||||
hs.kvs = make([]HeaderKeyValues, 0, len(h))
|
||||
}
|
||||
|
||||
kvs = hs.kvs[:0]
|
||||
for k, vv := range h {
|
||||
lock.RLock()
|
||||
if !exclude[k] {
|
||||
kvs = append(kvs, HeaderKeyValues{k, vv})
|
||||
}
|
||||
lock.RUnlock()
|
||||
}
|
||||
|
||||
hs.kvs = kvs
|
||||
// Clear any order left over from a previous SortedKeyValuesBy call on a
|
||||
// sorter that was returned to the pool; otherwise a default-order request
|
||||
// would inherit a stale custom order and sort non-lexicographically.
|
||||
hs.order = nil
|
||||
sort.Sort(hs)
|
||||
|
||||
return kvs, hs
|
||||
}
|
||||
|
||||
// SortedKeyValuesBy returns headers sorted by specified order
|
||||
func SortedKeyValuesBy(
|
||||
h http.Header,
|
||||
order map[string]int,
|
||||
exclude map[string]bool,
|
||||
) (kvs []HeaderKeyValues, hs *HeaderSorter) {
|
||||
hs = headerSorterPool.Get().(*HeaderSorter)
|
||||
if cap(hs.kvs) < len(h) {
|
||||
hs.kvs = make([]HeaderKeyValues, 0, len(h))
|
||||
}
|
||||
|
||||
kvs = hs.kvs[:0]
|
||||
for k, vv := range h {
|
||||
lock.RLock()
|
||||
if !exclude[k] {
|
||||
kvs = append(kvs, HeaderKeyValues{k, vv})
|
||||
}
|
||||
lock.RUnlock()
|
||||
}
|
||||
|
||||
hs.kvs = kvs
|
||||
hs.order = order
|
||||
sort.Sort(hs)
|
||||
|
||||
return kvs, hs
|
||||
}
|
||||
|
||||
// ReturnSorter returns a header sorter to the pool
|
||||
func ReturnSorter(hs *HeaderSorter) {
|
||||
headerSorterPool.Put(hs)
|
||||
}
|
||||
+48
@@ -0,0 +1,48 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"net"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// An addrList represents a list of network endpoint addresses.
|
||||
// Copy from [net.addrList] and change type from [net.Addr] to [net.IPAddr]
|
||||
type addrList []net.IPAddr
|
||||
|
||||
// isIPv4 reports whether addr contains an IPv4 address.
|
||||
func isIPv4(addr net.IPAddr) bool {
|
||||
return addr.IP.To4() != nil
|
||||
}
|
||||
|
||||
// isNotIPv4 reports whether addr does not contain an IPv4 address.
|
||||
func isNotIPv4(addr net.IPAddr) bool { return !isIPv4(addr) }
|
||||
|
||||
// forResolve returns the most appropriate address in address for
|
||||
// a call to ResolveTCPAddr, ResolveUDPAddr, or ResolveIPAddr.
|
||||
// IPv4 is preferred, unless addr contains an IPv6 literal.
|
||||
func (addrs addrList) forResolve(network, addr string) net.IPAddr {
|
||||
var want6 bool
|
||||
switch network {
|
||||
case "ip":
|
||||
// IPv6 literal (addr does NOT contain a port)
|
||||
want6 = strings.ContainsRune(addr, ':')
|
||||
case "tcp", "udp":
|
||||
// IPv6 literal. (addr contains a port, so look for '[')
|
||||
want6 = strings.ContainsRune(addr, '[')
|
||||
}
|
||||
if want6 {
|
||||
return addrs.first(isNotIPv4)
|
||||
}
|
||||
return addrs.first(isIPv4)
|
||||
}
|
||||
|
||||
// first returns the first address which satisfies strategy, or if
|
||||
// none do, then the first address of any kind.
|
||||
func (addrs addrList) first(strategy func(net.IPAddr) bool) net.IPAddr {
|
||||
for _, addr := range addrs {
|
||||
if strategy(addr) {
|
||||
return addr
|
||||
}
|
||||
}
|
||||
return addrs[0]
|
||||
}
|
||||
+11
@@ -0,0 +1,11 @@
|
||||
//go:build gomock || generate
|
||||
|
||||
package http3
|
||||
|
||||
//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -mock_names=TestClientConnInterface=MockClientConn -package http3 -destination mock_clientconn_test.go github.com/quic-go/quic-go/http3 TestClientConnInterface"
|
||||
type TestClientConnInterface = clientConn
|
||||
|
||||
//go:generate sh -c "go tool mockgen -typed -build_flags=\"-tags=gomock\" -mock_names=DatagramStream=MockDatagramStream -package http3 -destination mock_datagram_stream_test.go github.com/quic-go/quic-go/http3 DatagramStream"
|
||||
type DatagramStream = datagramStream
|
||||
|
||||
//go:generate sh -c "go tool mockgen -typed -package http3 -destination mock_quic_listener_test.go github.com/quic-go/quic-go/http3 QUICListener"
|
||||
+56
@@ -0,0 +1,56 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/enetx/http3/qlog"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
|
||||
"github.com/quic-go/qpack"
|
||||
)
|
||||
|
||||
func maybeQlogInvalidHeadersFrame(qlogger qlogwriter.Recorder, streamID quic.StreamID, l uint64) {
|
||||
if qlogger != nil {
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: streamID,
|
||||
Raw: qlog.RawInfo{PayloadLength: int(l)},
|
||||
Frame: qlog.Frame{Frame: qlog.HeadersFrame{}},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func qlogParsedHeadersFrame(qlogger qlogwriter.Recorder, streamID quic.StreamID, hf *headersFrame, hfs []qpack.HeaderField) {
|
||||
headerFields := make([]qlog.HeaderField, len(hfs))
|
||||
for i, hf := range hfs {
|
||||
headerFields[i] = qlog.HeaderField{
|
||||
Name: hf.Name,
|
||||
Value: hf.Value,
|
||||
}
|
||||
}
|
||||
qlogger.RecordEvent(qlog.FrameParsed{
|
||||
StreamID: streamID,
|
||||
Raw: qlog.RawInfo{
|
||||
Length: int(hf.Length) + hf.headerLen,
|
||||
PayloadLength: int(hf.Length),
|
||||
},
|
||||
Frame: qlog.Frame{Frame: qlog.HeadersFrame{
|
||||
HeaderFields: headerFields,
|
||||
}},
|
||||
})
|
||||
}
|
||||
|
||||
func qlogCreatedHeadersFrame(qlogger qlogwriter.Recorder, streamID quic.StreamID, length, payloadLength int, hfs []qlog.HeaderField) {
|
||||
headerFields := make([]qlog.HeaderField, len(hfs))
|
||||
for i, hf := range hfs {
|
||||
headerFields[i] = qlog.HeaderField{
|
||||
Name: hf.Name,
|
||||
Value: hf.Value,
|
||||
}
|
||||
}
|
||||
qlogger.RecordEvent(qlog.FrameCreated{
|
||||
StreamID: streamID,
|
||||
Raw: qlog.RawInfo{Length: length, PayloadLength: payloadLength},
|
||||
Frame: qlog.Frame{Frame: qlog.HeadersFrame{
|
||||
HeaderFields: headerFields,
|
||||
}},
|
||||
})
|
||||
}
|
||||
+138
@@ -0,0 +1,138 @@
|
||||
package qlog
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/qlogwriter/jsontext"
|
||||
)
|
||||
|
||||
type encoderHelper struct {
|
||||
enc *jsontext.Encoder
|
||||
err error
|
||||
}
|
||||
|
||||
func (h *encoderHelper) WriteToken(t jsontext.Token) {
|
||||
if h.err != nil {
|
||||
return
|
||||
}
|
||||
h.err = h.enc.WriteToken(t)
|
||||
}
|
||||
|
||||
type RawInfo struct {
|
||||
Length int // full packet length, including header and AEAD authentication tag
|
||||
PayloadLength int // length of the packet payload, excluding AEAD tag
|
||||
}
|
||||
|
||||
func (i RawInfo) HasValues() bool {
|
||||
return i.Length != 0 || i.PayloadLength != 0
|
||||
}
|
||||
|
||||
func (i RawInfo) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
if i.Length != 0 {
|
||||
h.WriteToken(jsontext.String("length"))
|
||||
h.WriteToken(jsontext.Uint(uint64(i.Length)))
|
||||
}
|
||||
if i.PayloadLength != 0 {
|
||||
h.WriteToken(jsontext.String("payload_length"))
|
||||
h.WriteToken(jsontext.Uint(uint64(i.PayloadLength)))
|
||||
}
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
type FrameParsed struct {
|
||||
StreamID quic.StreamID
|
||||
Raw RawInfo
|
||||
Frame Frame
|
||||
}
|
||||
|
||||
func (e FrameParsed) Name() string { return "http3:frame_parsed" }
|
||||
|
||||
func (e FrameParsed) Encode(enc *jsontext.Encoder, _ time.Time) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("stream_id"))
|
||||
h.WriteToken(jsontext.Uint(uint64(e.StreamID)))
|
||||
if e.Raw.HasValues() {
|
||||
h.WriteToken(jsontext.String("raw"))
|
||||
if err := e.Raw.encode(enc); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
h.WriteToken(jsontext.String("frame"))
|
||||
if err := e.Frame.encode(enc); err != nil {
|
||||
return err
|
||||
}
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
type FrameCreated struct {
|
||||
StreamID quic.StreamID
|
||||
Raw RawInfo
|
||||
Frame Frame
|
||||
}
|
||||
|
||||
func (e FrameCreated) Name() string { return "http3:frame_created" }
|
||||
|
||||
func (e FrameCreated) Encode(enc *jsontext.Encoder, _ time.Time) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("stream_id"))
|
||||
h.WriteToken(jsontext.Uint(uint64(e.StreamID)))
|
||||
if e.Raw.HasValues() {
|
||||
h.WriteToken(jsontext.String("raw"))
|
||||
if err := e.Raw.encode(enc); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
h.WriteToken(jsontext.String("frame"))
|
||||
if err := e.Frame.encode(enc); err != nil {
|
||||
return err
|
||||
}
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
type DatagramCreated struct {
|
||||
QuarterStreamID uint64
|
||||
Raw RawInfo
|
||||
}
|
||||
|
||||
func (e DatagramCreated) Name() string { return "http3:datagram_created" }
|
||||
|
||||
func (e DatagramCreated) Encode(enc *jsontext.Encoder, _ time.Time) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("quarter_stream_id"))
|
||||
h.WriteToken(jsontext.Uint(e.QuarterStreamID))
|
||||
h.WriteToken(jsontext.String("raw"))
|
||||
if err := e.Raw.encode(enc); err != nil {
|
||||
return err
|
||||
}
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
type DatagramParsed struct {
|
||||
QuarterStreamID uint64
|
||||
Raw RawInfo
|
||||
}
|
||||
|
||||
func (e DatagramParsed) Name() string { return "http3:datagram_parsed" }
|
||||
|
||||
func (e DatagramParsed) Encode(enc *jsontext.Encoder, _ time.Time) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("quarter_stream_id"))
|
||||
h.WriteToken(jsontext.Uint(e.QuarterStreamID))
|
||||
h.WriteToken(jsontext.String("raw"))
|
||||
if err := e.Raw.encode(enc); err != nil {
|
||||
return err
|
||||
}
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
+220
@@ -0,0 +1,220 @@
|
||||
package qlog
|
||||
|
||||
import (
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/qlogwriter/jsontext"
|
||||
)
|
||||
|
||||
// Frame represents an HTTP/3 frame.
|
||||
type Frame struct {
|
||||
Frame any
|
||||
}
|
||||
|
||||
func (f Frame) encode(enc *jsontext.Encoder) error {
|
||||
switch frame := f.Frame.(type) {
|
||||
case DataFrame:
|
||||
return frame.encode(enc)
|
||||
case HeadersFrame:
|
||||
return frame.encode(enc)
|
||||
case GoAwayFrame:
|
||||
return frame.encode(enc)
|
||||
case SettingsFrame:
|
||||
return frame.encode(enc)
|
||||
case PushPromiseFrame:
|
||||
return frame.encode(enc)
|
||||
case CancelPushFrame:
|
||||
return frame.encode(enc)
|
||||
case MaxPushIDFrame:
|
||||
return frame.encode(enc)
|
||||
case ReservedFrame:
|
||||
return frame.encode(enc)
|
||||
case UnknownFrame:
|
||||
return frame.encode(enc)
|
||||
}
|
||||
// This shouldn't happen if the code is correctly logging frames.
|
||||
// Write a null token to produce valid JSON.
|
||||
return enc.WriteToken(jsontext.Null)
|
||||
}
|
||||
|
||||
// A DataFrame is a DATA frame
|
||||
type DataFrame struct{}
|
||||
|
||||
func (f *DataFrame) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("frame_type"))
|
||||
h.WriteToken(jsontext.String("data"))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
type HeaderField struct {
|
||||
Name string
|
||||
Value string
|
||||
}
|
||||
|
||||
// A HeadersFrame is a HEADERS frame
|
||||
type HeadersFrame struct {
|
||||
HeaderFields []HeaderField
|
||||
}
|
||||
|
||||
func (f *HeadersFrame) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("frame_type"))
|
||||
h.WriteToken(jsontext.String("headers"))
|
||||
if len(f.HeaderFields) > 0 {
|
||||
h.WriteToken(jsontext.String("header_fields"))
|
||||
h.WriteToken(jsontext.BeginArray)
|
||||
for _, f := range f.HeaderFields {
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("name"))
|
||||
h.WriteToken(jsontext.String(f.Name))
|
||||
h.WriteToken(jsontext.String("value"))
|
||||
h.WriteToken(jsontext.String(f.Value))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
}
|
||||
h.WriteToken(jsontext.EndArray)
|
||||
}
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
// A GoAwayFrame is a GOAWAY frame
|
||||
type GoAwayFrame struct {
|
||||
StreamID quic.StreamID
|
||||
}
|
||||
|
||||
func (f *GoAwayFrame) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("frame_type"))
|
||||
h.WriteToken(jsontext.String("goaway"))
|
||||
h.WriteToken(jsontext.String("id"))
|
||||
h.WriteToken(jsontext.Uint(uint64(f.StreamID)))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
type SettingsFrame struct {
|
||||
MaxFieldSectionSize int64
|
||||
Datagram *bool
|
||||
ExtendedConnect *bool
|
||||
Other map[uint64]uint64
|
||||
}
|
||||
|
||||
func (f *SettingsFrame) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("frame_type"))
|
||||
h.WriteToken(jsontext.String("settings"))
|
||||
h.WriteToken(jsontext.String("settings"))
|
||||
h.WriteToken(jsontext.BeginArray)
|
||||
if f.MaxFieldSectionSize >= 0 {
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("name"))
|
||||
h.WriteToken(jsontext.String("settings_max_field_section_size"))
|
||||
h.WriteToken(jsontext.String("value"))
|
||||
h.WriteToken(jsontext.Uint(uint64(f.MaxFieldSectionSize)))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
}
|
||||
if f.Datagram != nil {
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("name"))
|
||||
h.WriteToken(jsontext.String("settings_h3_datagram"))
|
||||
h.WriteToken(jsontext.String("value"))
|
||||
h.WriteToken(jsontext.Bool(*f.Datagram))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
}
|
||||
if f.ExtendedConnect != nil {
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("name"))
|
||||
h.WriteToken(jsontext.String("settings_enable_connect_protocol"))
|
||||
h.WriteToken(jsontext.String("value"))
|
||||
h.WriteToken(jsontext.Bool(*f.ExtendedConnect))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
}
|
||||
if len(f.Other) > 0 {
|
||||
for k, v := range f.Other {
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("name"))
|
||||
h.WriteToken(jsontext.String("unknown"))
|
||||
h.WriteToken(jsontext.String("name_bytes"))
|
||||
h.WriteToken(jsontext.Uint(k))
|
||||
h.WriteToken(jsontext.String("value"))
|
||||
h.WriteToken(jsontext.Uint(v))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
}
|
||||
}
|
||||
h.WriteToken(jsontext.EndArray)
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
// A PushPromiseFrame is a PUSH_PROMISE frame
|
||||
type PushPromiseFrame struct{}
|
||||
|
||||
func (f *PushPromiseFrame) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("frame_type"))
|
||||
h.WriteToken(jsontext.String("push_promise"))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
// A CancelPushFrame is a CANCEL_PUSH frame
|
||||
type CancelPushFrame struct{}
|
||||
|
||||
func (f *CancelPushFrame) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("frame_type"))
|
||||
h.WriteToken(jsontext.String("cancel_push"))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
// A MaxPushIDFrame is a MAX_PUSH_ID frame
|
||||
type MaxPushIDFrame struct{}
|
||||
|
||||
func (f *MaxPushIDFrame) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("frame_type"))
|
||||
h.WriteToken(jsontext.String("max_push_id"))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
// A ReservedFrame is one of the reserved frame types
|
||||
type ReservedFrame struct {
|
||||
Type uint64
|
||||
}
|
||||
|
||||
func (f *ReservedFrame) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("frame_type"))
|
||||
h.WriteToken(jsontext.String("reserved"))
|
||||
h.WriteToken(jsontext.String("frame_type_bytes"))
|
||||
h.WriteToken(jsontext.Uint(f.Type))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
|
||||
// An UnknownFrame is an unknown frame type
|
||||
type UnknownFrame struct {
|
||||
Type uint64
|
||||
}
|
||||
|
||||
func (f *UnknownFrame) encode(enc *jsontext.Encoder) error {
|
||||
h := encoderHelper{enc: enc}
|
||||
h.WriteToken(jsontext.BeginObject)
|
||||
h.WriteToken(jsontext.String("frame_type"))
|
||||
h.WriteToken(jsontext.String("unknown"))
|
||||
h.WriteToken(jsontext.String("frame_type_bytes"))
|
||||
h.WriteToken(jsontext.Uint(f.Type))
|
||||
h.WriteToken(jsontext.EndObject)
|
||||
return h.err
|
||||
}
|
||||
+15
@@ -0,0 +1,15 @@
|
||||
package qlog
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/qlog"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
)
|
||||
|
||||
const EventSchema = "urn:ietf:params:qlog:events:http3-12"
|
||||
|
||||
func DefaultConnectionTracer(ctx context.Context, isClient bool, connID quic.ConnectionID) qlogwriter.Trace {
|
||||
return qlog.DefaultConnectionTracerWithSchemas(ctx, isClient, connID, []string{qlog.EventSchema, EventSchema})
|
||||
}
|
||||
+379
@@ -0,0 +1,379 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"github.com/enetx/http"
|
||||
"github.com/enetx/http/httptrace"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"golang.org/x/net/http/httpguts"
|
||||
"golang.org/x/net/http2/hpack"
|
||||
"golang.org/x/net/idna"
|
||||
|
||||
"github.com/quic-go/qpack"
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/enetx/http3/httpcommon"
|
||||
"github.com/enetx/http3/qlog"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
)
|
||||
|
||||
const bodyCopyBufferSize = 8 * 1024
|
||||
|
||||
type requestWriter struct {
|
||||
mutex sync.Mutex
|
||||
encoder *qpack.Encoder
|
||||
headerBuf *bytes.Buffer
|
||||
}
|
||||
|
||||
func newRequestWriter() *requestWriter {
|
||||
headerBuf := &bytes.Buffer{}
|
||||
encoder := qpack.NewEncoder(headerBuf)
|
||||
return &requestWriter{
|
||||
encoder: encoder,
|
||||
headerBuf: headerBuf,
|
||||
}
|
||||
}
|
||||
|
||||
func (w *requestWriter) WriteRequestHeader(wr io.Writer, req *http.Request, gzip bool, streamID quic.StreamID, qlogger qlogwriter.Recorder) error {
|
||||
buf := &bytes.Buffer{}
|
||||
if err := w.writeHeaders(buf, req, gzip, streamID, qlogger); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := wr.Write(buf.Bytes()); err != nil {
|
||||
return err
|
||||
}
|
||||
trace := httptrace.ContextClientTrace(req.Context())
|
||||
traceWroteHeaders(trace)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (w *requestWriter) writeHeaders(wr io.Writer, req *http.Request, gzip bool, streamID quic.StreamID, qlogger qlogwriter.Recorder) error {
|
||||
w.mutex.Lock()
|
||||
defer w.mutex.Unlock()
|
||||
defer w.encoder.Close()
|
||||
defer w.headerBuf.Reset()
|
||||
|
||||
var trailers string
|
||||
if len(req.Trailer) > 0 {
|
||||
keys := make([]string, 0, len(req.Trailer))
|
||||
for k := range req.Trailer {
|
||||
if httpguts.ValidTrailerHeader(k) {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
}
|
||||
trailers = strings.Join(keys, ", ")
|
||||
}
|
||||
|
||||
headerFields, err := w.encodeHeaders(req, gzip, trailers, actualContentLength(req), qlogger != nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
b := make([]byte, 0, 128)
|
||||
b = (&headersFrame{Length: uint64(w.headerBuf.Len())}).Append(b)
|
||||
if qlogger != nil {
|
||||
qlogCreatedHeadersFrame(qlogger, streamID, len(b)+w.headerBuf.Len(), w.headerBuf.Len(), headerFields)
|
||||
}
|
||||
if _, err := wr.Write(b); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = wr.Write(w.headerBuf.Bytes())
|
||||
return err
|
||||
}
|
||||
|
||||
func isExtendedConnectRequest(req *http.Request) bool {
|
||||
return req.Method == http.MethodConnect && req.Proto != "" && req.Proto != "HTTP/1.1"
|
||||
}
|
||||
|
||||
// copied from net/transport.go
|
||||
// Modified to support Extended CONNECT:
|
||||
// Contrary to what the godoc for the http.Request says,
|
||||
// we do respect the Proto field if the method is CONNECT.
|
||||
//
|
||||
// The returned header fields are only set if doQlog is true.
|
||||
func (w *requestWriter) encodeHeaders(req *http.Request, addGzipHeader bool, trailers string, contentLength int64, doQlog bool) ([]qlog.HeaderField, error) {
|
||||
host := req.Host
|
||||
if host == "" {
|
||||
host = req.URL.Host
|
||||
}
|
||||
host, err := httpguts.PunycodeHostPort(host)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !httpguts.ValidHostHeader(host) {
|
||||
return nil, errors.New("http3: invalid Host header")
|
||||
}
|
||||
|
||||
// http.NewRequest sets this field to HTTP/1.1
|
||||
isExtendedConnect := isExtendedConnectRequest(req)
|
||||
if isExtendedConnect && !validExtendedConnectProtocol(req.Proto) {
|
||||
return nil, fmt.Errorf("invalid request :protocol %q", req.Proto)
|
||||
}
|
||||
|
||||
var path string
|
||||
if req.Method != http.MethodConnect || isExtendedConnect {
|
||||
path = req.URL.RequestURI()
|
||||
if !validPseudoPath(path) {
|
||||
orig := path
|
||||
path = strings.TrimPrefix(path, req.URL.Scheme+"://"+host)
|
||||
if !validPseudoPath(path) {
|
||||
if req.URL.Opaque != "" {
|
||||
return nil, fmt.Errorf("invalid request :path %q from URL.Opaque = %q", orig, req.URL.Opaque)
|
||||
} else {
|
||||
return nil, fmt.Errorf("invalid request :path %q", orig)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check for any invalid headers and return an error before we
|
||||
// potentially pollute our hpack state. (We want to be able to
|
||||
// continue to reuse the hpack encoder for future requests)
|
||||
for k, vv := range req.Header {
|
||||
// Skip validation for special header order keys
|
||||
if k == httpcommon.HeaderOrderKey || k == httpcommon.PHeaderOrderKey {
|
||||
continue
|
||||
}
|
||||
if !httpguts.ValidHeaderFieldName(k) {
|
||||
return nil, fmt.Errorf("invalid HTTP header name %q", k)
|
||||
}
|
||||
for _, v := range vv {
|
||||
if !httpguts.ValidHeaderFieldValue(v) {
|
||||
return nil, fmt.Errorf("invalid HTTP header value %q for header %q", v, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
enumerateHeaders := func(f func(name, value string)) {
|
||||
// Handle pseudo-headers with order support
|
||||
pHeaderOrder, hasPHeaderOrder := req.Header[httpcommon.PHeaderOrderKey]
|
||||
|
||||
if hasPHeaderOrder {
|
||||
// Follow pseudo header order
|
||||
for _, p := range pHeaderOrder {
|
||||
switch p {
|
||||
case ":authority":
|
||||
f(":authority", host)
|
||||
case ":method":
|
||||
f(":method", req.Method)
|
||||
case ":path":
|
||||
if req.Method != http.MethodConnect || isExtendedConnect {
|
||||
f(":path", path)
|
||||
}
|
||||
case ":scheme":
|
||||
if req.Method != http.MethodConnect || isExtendedConnect {
|
||||
f(":scheme", req.URL.Scheme)
|
||||
}
|
||||
case ":protocol":
|
||||
if isExtendedConnect {
|
||||
f(":protocol", req.Proto)
|
||||
}
|
||||
default:
|
||||
continue
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// Default pseudo-header order
|
||||
f(":authority", host)
|
||||
f(":method", req.Method)
|
||||
if req.Method != http.MethodConnect || isExtendedConnect {
|
||||
f(":path", path)
|
||||
f(":scheme", req.URL.Scheme)
|
||||
}
|
||||
if isExtendedConnect {
|
||||
f(":protocol", req.Proto)
|
||||
}
|
||||
}
|
||||
|
||||
if trailers != "" {
|
||||
f("trailer", trailers)
|
||||
}
|
||||
|
||||
// Handle regular headers with order support
|
||||
exclude := make(map[string]bool)
|
||||
exclude[httpcommon.HeaderOrderKey] = true
|
||||
exclude[httpcommon.PHeaderOrderKey] = true
|
||||
|
||||
var kvs []httpcommon.HeaderKeyValues
|
||||
var sorter *httpcommon.HeaderSorter
|
||||
|
||||
if headerOrder, hasHeaderOrder := req.Header[httpcommon.HeaderOrderKey]; hasHeaderOrder {
|
||||
order := make(map[string]int)
|
||||
for i, v := range headerOrder {
|
||||
order[v] = i
|
||||
}
|
||||
kvs, sorter = httpcommon.SortedKeyValuesBy(req.Header, order, exclude)
|
||||
} else {
|
||||
kvs, sorter = httpcommon.SortedKeyValues(req.Header, exclude)
|
||||
}
|
||||
|
||||
var didUA bool
|
||||
for _, kv := range kvs {
|
||||
k := kv.Key
|
||||
vv := kv.Values
|
||||
|
||||
if strings.EqualFold(k, "host") || strings.EqualFold(k, "content-length") {
|
||||
// Host is :authority, already sent.
|
||||
// Content-Length is automatic, set below.
|
||||
continue
|
||||
} else if strings.EqualFold(k, "connection") || strings.EqualFold(k, "proxy-connection") ||
|
||||
strings.EqualFold(k, "transfer-encoding") || strings.EqualFold(k, "upgrade") ||
|
||||
strings.EqualFold(k, "keep-alive") {
|
||||
// Per 8.1.2.2 Connection-Specific Header
|
||||
// Fields, don't send connection-specific
|
||||
// fields. We have already checked if any
|
||||
// are error-worthy so just ignore the rest.
|
||||
continue
|
||||
} else if strings.EqualFold(k, "user-agent") {
|
||||
// Match Go's http1 behavior: at most one
|
||||
// User-Agent. If set to nil or empty string,
|
||||
// then omit it. Otherwise if not mentioned,
|
||||
// include the default (below).
|
||||
didUA = true
|
||||
if len(vv) < 1 {
|
||||
continue
|
||||
}
|
||||
vv = vv[:1]
|
||||
if vv[0] == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
for _, v := range vv {
|
||||
f(k, v)
|
||||
}
|
||||
}
|
||||
|
||||
httpcommon.ReturnSorter(sorter)
|
||||
|
||||
if shouldSendReqContentLength(req.Method, contentLength) {
|
||||
f("content-length", strconv.FormatInt(contentLength, 10))
|
||||
}
|
||||
if addGzipHeader {
|
||||
f("accept-encoding", "gzip")
|
||||
}
|
||||
if !didUA {
|
||||
f("user-agent", defaultUserAgent)
|
||||
}
|
||||
}
|
||||
|
||||
// Do a first pass over the headers counting bytes to ensure
|
||||
// we don't exceed cc.peerMaxHeaderListSize. This is done as a
|
||||
// separate pass before encoding the headers to prevent
|
||||
// modifying the hpack state.
|
||||
hlSize := uint64(0)
|
||||
enumerateHeaders(func(name, value string) {
|
||||
hf := hpack.HeaderField{Name: name, Value: value}
|
||||
hlSize += uint64(hf.Size())
|
||||
})
|
||||
|
||||
// TODO: check maximum header list size
|
||||
// if hlSize > cc.peerMaxHeaderListSize {
|
||||
// return errRequestHeaderListSize
|
||||
// }
|
||||
|
||||
trace := httptrace.ContextClientTrace(req.Context())
|
||||
traceHeaders := traceHasWroteHeaderField(trace)
|
||||
|
||||
// Header list size is ok. Write the headers.
|
||||
var headerFields []qlog.HeaderField
|
||||
if doQlog {
|
||||
headerFields = make([]qlog.HeaderField, 0, len(req.Header))
|
||||
}
|
||||
enumerateHeaders(func(name, value string) {
|
||||
name = strings.ToLower(name)
|
||||
w.encoder.WriteField(qpack.HeaderField{Name: name, Value: value})
|
||||
if traceHeaders {
|
||||
traceWroteHeaderField(trace, name, value)
|
||||
}
|
||||
if doQlog {
|
||||
headerFields = append(headerFields, qlog.HeaderField{Name: name, Value: value})
|
||||
}
|
||||
})
|
||||
|
||||
return headerFields, nil
|
||||
}
|
||||
|
||||
// authorityAddr returns a given authority (a host/IP, or host:port / ip:port)
|
||||
// and returns a host:port. The port 443 is added if needed.
|
||||
func authorityAddr(authority string) (addr string) {
|
||||
host, port, err := net.SplitHostPort(authority)
|
||||
if err != nil { // authority didn't have a port
|
||||
port = "443"
|
||||
host = authority
|
||||
}
|
||||
if a, err := idna.ToASCII(host); err == nil {
|
||||
host = a
|
||||
}
|
||||
// IPv6 address literal, without a port:
|
||||
if strings.HasPrefix(host, "[") && strings.HasSuffix(host, "]") {
|
||||
return host + ":" + port
|
||||
}
|
||||
return net.JoinHostPort(host, port)
|
||||
}
|
||||
|
||||
// validPseudoPath reports whether v is a valid :path pseudo-header
|
||||
// value. It must be either:
|
||||
//
|
||||
// *) a non-empty string starting with '/'
|
||||
// *) the string '*', for OPTIONS requests.
|
||||
//
|
||||
// For now this is only used a quick check for deciding when to clean
|
||||
// up Opaque URLs before sending requests from the Transport.
|
||||
// See golang.org/issue/16847
|
||||
//
|
||||
// We used to enforce that the path also didn't start with "//", but
|
||||
// Google's GFE accepts such paths and Chrome sends them, so ignore
|
||||
// that part of the spec. See golang.org/issue/19103.
|
||||
func validPseudoPath(v string) bool {
|
||||
return (len(v) > 0 && v[0] == '/') || v == "*"
|
||||
}
|
||||
|
||||
// actualContentLength returns a sanitized version of
|
||||
// req.ContentLength, where 0 actually means zero (not unknown) and -1
|
||||
// means unknown.
|
||||
func actualContentLength(req *http.Request) int64 {
|
||||
if req.Body == nil {
|
||||
return 0
|
||||
}
|
||||
if req.ContentLength != 0 {
|
||||
return req.ContentLength
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
// shouldSendReqContentLength reports whether the http2.Transport should send
|
||||
// a "content-length" request header. This logic is basically a copy of the net/http
|
||||
// transferWriter.shouldSendContentLength.
|
||||
// The contentLength is the corrected contentLength (so 0 means actually 0, not unknown).
|
||||
// -1 means unknown.
|
||||
func shouldSendReqContentLength(method string, contentLength int64) bool {
|
||||
if contentLength > 0 {
|
||||
return true
|
||||
}
|
||||
if contentLength < 0 {
|
||||
return false
|
||||
}
|
||||
// For zero bodies, whether we send a content-length depends on the method.
|
||||
// It also kinda doesn't matter for http2 either way, with END_STREAM.
|
||||
switch method {
|
||||
case "POST", "PUT", "PATCH":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// WriteRequestTrailer writes HTTP trailers to the stream.
|
||||
// It should be called after the request body has been fully written.
|
||||
func (w *requestWriter) WriteRequestTrailer(wr io.Writer, req *http.Request, streamID quic.StreamID, qlogger qlogwriter.Recorder) error {
|
||||
_, err := writeTrailers(wr, req.Trailer, streamID, qlogger)
|
||||
return err
|
||||
}
|
||||
+372
@@ -0,0 +1,372 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"github.com/enetx/http"
|
||||
"net/textproto"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/qpack"
|
||||
"github.com/enetx/http3/qlog"
|
||||
|
||||
"golang.org/x/net/http/httpguts"
|
||||
)
|
||||
|
||||
// The HTTPStreamer allows taking over a HTTP/3 stream. The interface is implemented by the http.ResponseWriter.
|
||||
// When a stream is taken over, it's the caller's responsibility to close the stream.
|
||||
type HTTPStreamer interface {
|
||||
HTTPStream() *Stream
|
||||
}
|
||||
|
||||
const maxSmallResponseSize = 4096
|
||||
|
||||
type responseWriter struct {
|
||||
str *Stream
|
||||
|
||||
conn *rawConn
|
||||
header http.Header
|
||||
trailers map[string]struct{}
|
||||
buf []byte
|
||||
status int // status code passed to WriteHeader
|
||||
|
||||
// for responses smaller than maxSmallResponseSize, we buffer calls to Write,
|
||||
// and automatically add the Content-Length header
|
||||
smallResponseBuf []byte
|
||||
|
||||
contentLen int64 // if handler set valid Content-Length header
|
||||
numWritten int64 // bytes written
|
||||
headerComplete bool // set once WriteHeader is called with a status code >= 200
|
||||
headerWritten bool // set once the response header has been serialized to the stream
|
||||
isHead bool
|
||||
trailerWritten bool // set once the response trailers has been serialized to the stream
|
||||
|
||||
hijacked bool // set on HTTPStream is called
|
||||
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
var (
|
||||
_ http.ResponseWriter = &responseWriter{}
|
||||
_ http.Flusher = &responseWriter{}
|
||||
_ Settingser = &responseWriter{}
|
||||
_ HTTPStreamer = &responseWriter{}
|
||||
// make sure that we implement (some of the) methods used by the http.ResponseController
|
||||
_ interface {
|
||||
SetReadDeadline(time.Time) error
|
||||
SetWriteDeadline(time.Time) error
|
||||
Flush()
|
||||
FlushError() error
|
||||
} = &responseWriter{}
|
||||
)
|
||||
|
||||
func newResponseWriter(str *Stream, conn *rawConn, isHead bool, logger *slog.Logger) *responseWriter {
|
||||
return &responseWriter{
|
||||
str: str,
|
||||
conn: conn,
|
||||
header: http.Header{},
|
||||
buf: make([]byte, frameHeaderLen),
|
||||
isHead: isHead,
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
func (w *responseWriter) Header() http.Header {
|
||||
return w.header
|
||||
}
|
||||
|
||||
func (w *responseWriter) WriteHeader(status int) {
|
||||
if w.headerComplete {
|
||||
return
|
||||
}
|
||||
|
||||
// http status must be 3 digits
|
||||
if status < 100 || status > 999 {
|
||||
panic(fmt.Sprintf("invalid WriteHeader code %v", status))
|
||||
}
|
||||
w.status = status
|
||||
|
||||
// immediately write 1xx headers
|
||||
if status < 200 {
|
||||
w.writeHeader(status)
|
||||
return
|
||||
}
|
||||
|
||||
// We're done with headers once we write a status >= 200.
|
||||
w.headerComplete = true
|
||||
// Add Date header.
|
||||
// This is what the standard library does.
|
||||
// Can be disabled by setting the Date header to nil.
|
||||
if _, ok := w.header["Date"]; !ok {
|
||||
w.header.Set("Date", time.Now().UTC().Format(http.TimeFormat))
|
||||
}
|
||||
// Content-Length checking
|
||||
// use ParseUint instead of ParseInt, as negative values are invalid
|
||||
if clen := w.header.Get("Content-Length"); clen != "" {
|
||||
if cl, err := strconv.ParseUint(clen, 10, 63); err == nil {
|
||||
w.contentLen = int64(cl)
|
||||
} else {
|
||||
// emit a warning for malformed Content-Length and remove it
|
||||
logger := w.logger
|
||||
if logger == nil {
|
||||
logger = slog.Default()
|
||||
}
|
||||
logger.Error("Malformed Content-Length", "value", clen)
|
||||
w.header.Del("Content-Length")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (w *responseWriter) sniffContentType(p []byte) {
|
||||
// If no content type, apply sniffing algorithm to body.
|
||||
// We can't use `w.header.Get` here since if the Content-Type was set to nil, we shouldn't do sniffing.
|
||||
_, haveType := w.header["Content-Type"]
|
||||
|
||||
// If the Content-Encoding was set and is non-blank, we shouldn't sniff the body.
|
||||
hasCE := w.header.Get("Content-Encoding") != ""
|
||||
if !hasCE && !haveType && len(p) > 0 {
|
||||
w.header.Set("Content-Type", http.DetectContentType(p))
|
||||
}
|
||||
}
|
||||
|
||||
func (w *responseWriter) Write(p []byte) (int, error) {
|
||||
bodyAllowed := bodyAllowedForStatus(w.status)
|
||||
if !w.headerComplete {
|
||||
w.sniffContentType(p)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
bodyAllowed = true
|
||||
}
|
||||
if !bodyAllowed {
|
||||
return 0, http.ErrBodyNotAllowed
|
||||
}
|
||||
|
||||
w.numWritten += int64(len(p))
|
||||
if w.contentLen != 0 && w.numWritten > w.contentLen {
|
||||
return 0, http.ErrContentLength
|
||||
}
|
||||
|
||||
if w.isHead {
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
if !w.headerWritten {
|
||||
// Buffer small responses.
|
||||
// This allows us to automatically set the Content-Length field.
|
||||
if len(w.smallResponseBuf)+len(p) < maxSmallResponseSize {
|
||||
w.smallResponseBuf = append(w.smallResponseBuf, p...)
|
||||
return len(p), nil
|
||||
}
|
||||
}
|
||||
return w.doWrite(p)
|
||||
}
|
||||
|
||||
func (w *responseWriter) doWrite(p []byte) (int, error) {
|
||||
if !w.headerWritten {
|
||||
w.sniffContentType(w.smallResponseBuf)
|
||||
if err := w.writeHeader(w.status); err != nil {
|
||||
return 0, maybeReplaceError(err)
|
||||
}
|
||||
w.headerWritten = true
|
||||
}
|
||||
|
||||
l := uint64(len(w.smallResponseBuf) + len(p))
|
||||
if l == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
df := &dataFrame{Length: l}
|
||||
w.buf = w.buf[:0]
|
||||
w.buf = df.Append(w.buf)
|
||||
if w.str.qlogger != nil {
|
||||
w.str.qlogger.RecordEvent(qlog.FrameCreated{
|
||||
StreamID: w.str.StreamID(),
|
||||
Raw: qlog.RawInfo{Length: len(w.buf) + int(l), PayloadLength: int(l)},
|
||||
Frame: qlog.Frame{Frame: qlog.DataFrame{}},
|
||||
})
|
||||
}
|
||||
if _, err := w.str.writeUnframed(w.buf); err != nil {
|
||||
return 0, maybeReplaceError(err)
|
||||
}
|
||||
if len(w.smallResponseBuf) > 0 {
|
||||
if _, err := w.str.writeUnframed(w.smallResponseBuf); err != nil {
|
||||
return 0, maybeReplaceError(err)
|
||||
}
|
||||
w.smallResponseBuf = nil
|
||||
}
|
||||
var n int
|
||||
if len(p) > 0 {
|
||||
var err error
|
||||
n, err = w.str.writeUnframed(p)
|
||||
if err != nil {
|
||||
return n, maybeReplaceError(err)
|
||||
}
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (w *responseWriter) writeHeader(status int) error {
|
||||
var headerFields []qlog.HeaderField // only used for qlog
|
||||
var headers bytes.Buffer
|
||||
enc := qpack.NewEncoder(&headers)
|
||||
if err := enc.WriteField(qpack.HeaderField{Name: ":status", Value: strconv.Itoa(status)}); err != nil {
|
||||
return err
|
||||
}
|
||||
if w.str.qlogger != nil {
|
||||
headerFields = append(headerFields, qlog.HeaderField{Name: ":status", Value: strconv.Itoa(status)})
|
||||
}
|
||||
|
||||
// Handle trailer fields
|
||||
if vals, ok := w.header["Trailer"]; ok {
|
||||
for _, val := range vals {
|
||||
for trailer := range strings.SplitSeq(val, ",") {
|
||||
// We need to convert to the canonical header key value here because this will be called when using
|
||||
// headers.Add or headers.Set.
|
||||
trailer = textproto.CanonicalMIMEHeaderKey(strings.TrimSpace(trailer))
|
||||
w.declareTrailer(trailer)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for k, v := range w.header {
|
||||
if _, excluded := w.trailers[k]; excluded {
|
||||
continue
|
||||
}
|
||||
// Ignore "Trailer:" prefixed headers
|
||||
if strings.HasPrefix(k, http.TrailerPrefix) {
|
||||
continue
|
||||
}
|
||||
for index := range v {
|
||||
name := strings.ToLower(k)
|
||||
value := v[index]
|
||||
if err := enc.WriteField(qpack.HeaderField{Name: name, Value: value}); err != nil {
|
||||
return err
|
||||
}
|
||||
if w.str.qlogger != nil {
|
||||
headerFields = append(headerFields, qlog.HeaderField{Name: name, Value: value})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
buf := make([]byte, 0, frameHeaderLen+headers.Len())
|
||||
buf = (&headersFrame{Length: uint64(headers.Len())}).Append(buf)
|
||||
buf = append(buf, headers.Bytes()...)
|
||||
|
||||
if w.str.qlogger != nil {
|
||||
qlogCreatedHeadersFrame(w.str.qlogger, w.str.StreamID(), len(buf), headers.Len(), headerFields)
|
||||
}
|
||||
|
||||
_, err := w.str.writeUnframed(buf)
|
||||
return err
|
||||
}
|
||||
|
||||
func (w *responseWriter) FlushError() error {
|
||||
if !w.headerComplete {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}
|
||||
_, err := w.doWrite(nil)
|
||||
return err
|
||||
}
|
||||
|
||||
func (w *responseWriter) flushTrailers() {
|
||||
if w.trailerWritten {
|
||||
return
|
||||
}
|
||||
if err := w.writeTrailers(); err != nil {
|
||||
if w.logger != nil {
|
||||
w.logger.Debug("could not write trailers", "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (w *responseWriter) Flush() {
|
||||
if err := w.FlushError(); err != nil {
|
||||
if w.logger != nil {
|
||||
w.logger.Debug("could not flush to stream", "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// declareTrailer adds a trailer to the trailer list, while also validating that the trailer has a
|
||||
// valid name.
|
||||
func (w *responseWriter) declareTrailer(k string) {
|
||||
if !httpguts.ValidTrailerHeader(k) {
|
||||
// Forbidden by RFC 9110, section 6.5.1.
|
||||
if w.logger != nil {
|
||||
w.logger.Debug("ignoring invalid trailer", slog.String("header", k))
|
||||
}
|
||||
return
|
||||
}
|
||||
if w.trailers == nil {
|
||||
w.trailers = make(map[string]struct{})
|
||||
}
|
||||
w.trailers[k] = struct{}{}
|
||||
}
|
||||
|
||||
// writeTrailers will write trailers to the stream if there are any.
|
||||
func (w *responseWriter) writeTrailers() error {
|
||||
// promote headers added via "Trailer:" convention as trailers, these can be added after
|
||||
// streaming the status/headers have been written.
|
||||
for k := range w.header {
|
||||
if strings.HasPrefix(k, http.TrailerPrefix) {
|
||||
w.declareTrailer(k)
|
||||
}
|
||||
}
|
||||
|
||||
if len(w.trailers) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
trailers := make(http.Header, len(w.trailers))
|
||||
for trailer := range w.trailers {
|
||||
if vals, ok := w.header[trailer]; ok {
|
||||
trailers[strings.TrimPrefix(trailer, http.TrailerPrefix)] = vals
|
||||
}
|
||||
}
|
||||
|
||||
written, err := writeTrailers(w.str.datagramStream, trailers, w.str.StreamID(), w.str.qlogger)
|
||||
if written {
|
||||
w.trailerWritten = true
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (w *responseWriter) HTTPStream() *Stream {
|
||||
w.hijacked = true
|
||||
w.Flush()
|
||||
return w.str
|
||||
}
|
||||
|
||||
func (w *responseWriter) wasStreamHijacked() bool { return w.hijacked }
|
||||
|
||||
func (w *responseWriter) ReceivedSettings() <-chan struct{} {
|
||||
return w.conn.ReceivedSettings()
|
||||
}
|
||||
|
||||
func (w *responseWriter) Settings() *Settings {
|
||||
return w.conn.Settings()
|
||||
}
|
||||
|
||||
func (w *responseWriter) SetReadDeadline(deadline time.Time) error {
|
||||
return w.str.SetReadDeadline(deadline)
|
||||
}
|
||||
|
||||
func (w *responseWriter) SetWriteDeadline(deadline time.Time) error {
|
||||
return w.str.SetWriteDeadline(deadline)
|
||||
}
|
||||
|
||||
// copied from http2/http2.go
|
||||
// bodyAllowedForStatus reports whether a given response status code
|
||||
// permits a body. See RFC 2616, section 4.4.
|
||||
func bodyAllowedForStatus(status int) bool {
|
||||
switch {
|
||||
case status >= 100 && status <= 199:
|
||||
return false
|
||||
case status == http.StatusNoContent:
|
||||
return false
|
||||
case status == http.StatusNotModified:
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
+740
@@ -0,0 +1,740 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/http"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/http3/qlog"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
)
|
||||
|
||||
// NextProtoH3 is the ALPN protocol negotiated during the TLS handshake, for QUIC v1 and v2.
|
||||
const NextProtoH3 = "h3"
|
||||
|
||||
// StreamType is the stream type of a unidirectional stream.
|
||||
type StreamType uint64
|
||||
|
||||
const (
|
||||
streamTypeControlStream = 0
|
||||
streamTypePushStream = 1
|
||||
streamTypeQPACKEncoderStream = 2
|
||||
streamTypeQPACKDecoderStream = 3
|
||||
)
|
||||
|
||||
// A QUICListener listens for incoming QUIC connections.
|
||||
type QUICListener interface {
|
||||
Accept(context.Context) (*quic.Conn, error)
|
||||
Addr() net.Addr
|
||||
io.Closer
|
||||
}
|
||||
|
||||
var _ QUICListener = &quic.EarlyListener{}
|
||||
|
||||
// ConfigureTLSConfig creates a new tls.Config which can be used
|
||||
// to create a quic.Listener meant for serving HTTP/3.
|
||||
func ConfigureTLSConfig(tlsConf *tls.Config) *tls.Config {
|
||||
// Workaround for https://github.com/golang/go/issues/60506.
|
||||
// This initializes the session tickets _before_ cloning the config.
|
||||
_, _ = tlsConf.DecryptTicket(nil, tls.ConnectionState{})
|
||||
config := tlsConf.Clone()
|
||||
config.NextProtos = []string{NextProtoH3}
|
||||
if gfc := config.GetConfigForClient; gfc != nil {
|
||||
config.GetConfigForClient = func(ch *tls.ClientHelloInfo) (*tls.Config, error) {
|
||||
conf, err := gfc(ch)
|
||||
if conf == nil || err != nil {
|
||||
return conf, err
|
||||
}
|
||||
return ConfigureTLSConfig(conf), nil
|
||||
}
|
||||
}
|
||||
return config
|
||||
}
|
||||
|
||||
// contextKey is a value for use with context.WithValue. It's used as
|
||||
// a pointer so it fits in an interface{} without allocation.
|
||||
type contextKey struct {
|
||||
name string
|
||||
}
|
||||
|
||||
func (k *contextKey) String() string { return "quic-go/http3 context value " + k.name }
|
||||
|
||||
// ServerContextKey is a context key. It can be used in HTTP
|
||||
// handlers with Context.Value to access the server that
|
||||
// started the handler. The associated value will be of
|
||||
// type *http3.Server.
|
||||
var ServerContextKey = &contextKey{"http3-server"}
|
||||
|
||||
// RemoteAddrContextKey is a context key. It can be used in
|
||||
// HTTP handlers with Context.Value to access the remote
|
||||
// address of the connection. The associated value will be of
|
||||
// type net.Addr.
|
||||
//
|
||||
// Use this value instead of [http.Request.RemoteAddr] if you
|
||||
// require access to the remote address of the connection rather
|
||||
// than its string representation.
|
||||
var RemoteAddrContextKey = &contextKey{"remote-addr"}
|
||||
|
||||
// listener contains info about specific listener added with addListener
|
||||
type listener struct {
|
||||
ln *QUICListener
|
||||
port int // 0 means that no info about port is available
|
||||
|
||||
// if this listener was constructed by the application, it won't be closed when the server is closed
|
||||
createdLocally bool
|
||||
}
|
||||
|
||||
// Server is a HTTP/3 server.
|
||||
type Server struct {
|
||||
// Addr optionally specifies the UDP address for the server to listen on,
|
||||
// in the form "host:port".
|
||||
//
|
||||
// When used by ListenAndServe and ListenAndServeTLS methods, if empty,
|
||||
// ":https" (port 443) is used. See net.Dial for details of the address
|
||||
// format.
|
||||
//
|
||||
// Otherwise, if Port is not set and underlying QUIC listeners do not
|
||||
// have valid port numbers, the port part is used in Alt-Svc headers set
|
||||
// with SetQUICHeaders.
|
||||
Addr string
|
||||
|
||||
// Port is used in Alt-Svc response headers set with SetQUICHeaders. If
|
||||
// needed Port can be manually set when the Server is created.
|
||||
//
|
||||
// This is useful when a Layer 4 firewall is redirecting UDP traffic and
|
||||
// clients must use a port different from the port the Server is
|
||||
// listening on.
|
||||
Port int
|
||||
|
||||
// TLSConfig provides a TLS configuration for use by server. It must be
|
||||
// set for ListenAndServe and Serve methods.
|
||||
TLSConfig *tls.Config
|
||||
|
||||
// QUICConfig provides the parameters for QUIC connection created with Serve.
|
||||
// If nil, it uses reasonable default values.
|
||||
//
|
||||
// Configured versions are also used in Alt-Svc response header set with SetQUICHeaders.
|
||||
QUICConfig *quic.Config
|
||||
|
||||
// Handler is the HTTP request handler to use. If not set, defaults to
|
||||
// http.NotFound.
|
||||
Handler http.Handler
|
||||
|
||||
// EnableDatagrams enables support for HTTP/3 datagrams (RFC 9297).
|
||||
// If set to true, QUICConfig.EnableDatagrams will be set.
|
||||
EnableDatagrams bool
|
||||
|
||||
// MaxHeaderBytes controls the maximum number of bytes the server will
|
||||
// read parsing the request HEADERS frame. It does not limit the size of
|
||||
// the request body. If zero or negative, http.DefaultMaxHeaderBytes is
|
||||
// used.
|
||||
MaxHeaderBytes int
|
||||
|
||||
// AdditionalSettings specifies additional HTTP/3 settings.
|
||||
// It is invalid to specify any settings defined by RFC 9114 (HTTP/3) and RFC 9297 (HTTP Datagrams).
|
||||
AdditionalSettings g.MapOrd[uint64, uint64]
|
||||
|
||||
// IdleTimeout specifies how long until idle clients connection should be
|
||||
// closed. Idle refers only to the HTTP/3 layer, activity at the QUIC layer
|
||||
// like PING frames are not considered.
|
||||
// If zero or negative, there is no timeout.
|
||||
IdleTimeout time.Duration
|
||||
|
||||
// ConnContext optionally specifies a function that modifies the context used for a new connection c.
|
||||
// The provided ctx has a ServerContextKey value.
|
||||
ConnContext func(ctx context.Context, c *quic.Conn) context.Context
|
||||
|
||||
Logger *slog.Logger
|
||||
|
||||
mutex sync.RWMutex
|
||||
listeners []listener
|
||||
|
||||
closed bool
|
||||
closeCtx context.Context // canceled when the server is closed
|
||||
closeCancel context.CancelFunc // cancels the closeCtx
|
||||
graceCtx context.Context // canceled when the server is closed or gracefully closed
|
||||
graceCancel context.CancelFunc // cancels the graceCtx
|
||||
connCount atomic.Int64
|
||||
connHandlingDone chan struct{}
|
||||
|
||||
altSvcHeader string
|
||||
}
|
||||
|
||||
// ListenAndServe listens on the UDP address s.Addr and calls s.Handler to handle HTTP/3 requests on incoming connections.
|
||||
//
|
||||
// If s.Addr is blank, ":https" is used.
|
||||
func (s *Server) ListenAndServe() error {
|
||||
ln, err := s.setupListenerForConn(s.TLSConfig, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer s.removeListener(ln)
|
||||
|
||||
return s.serveListener(*ln)
|
||||
}
|
||||
|
||||
// ListenAndServeTLS listens on the UDP address s.Addr and calls s.Handler to handle HTTP/3 requests on incoming connections.
|
||||
//
|
||||
// If s.Addr is blank, ":https" is used.
|
||||
func (s *Server) ListenAndServeTLS(certFile, keyFile string) error {
|
||||
var err error
|
||||
certs := make([]tls.Certificate, 1)
|
||||
certs[0], err = tls.LoadX509KeyPair(certFile, keyFile)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// We currently only use the cert-related stuff from tls.Config,
|
||||
// so we don't need to make a full copy.
|
||||
ln, err := s.setupListenerForConn(&tls.Config{Certificates: certs}, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer s.removeListener(ln)
|
||||
|
||||
return s.serveListener(*ln)
|
||||
}
|
||||
|
||||
// Serve an existing UDP connection.
|
||||
// It is possible to reuse the same connection for outgoing connections.
|
||||
// Closing the server does not close the connection.
|
||||
func (s *Server) Serve(conn net.PacketConn) error {
|
||||
ln, err := s.setupListenerForConn(s.TLSConfig, conn)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer s.removeListener(ln)
|
||||
|
||||
return s.serveListener(*ln)
|
||||
}
|
||||
|
||||
// init initializes the contexts used for shutting down the server.
|
||||
// It must be called with the mutex held.
|
||||
func (s *Server) init() {
|
||||
if s.closeCtx == nil {
|
||||
s.closeCtx, s.closeCancel = context.WithCancel(context.Background())
|
||||
s.graceCtx, s.graceCancel = context.WithCancel(s.closeCtx)
|
||||
}
|
||||
s.connHandlingDone = make(chan struct{}, 1)
|
||||
}
|
||||
|
||||
func (s *Server) decreaseConnCount() {
|
||||
if s.connCount.Add(-1) == 0 && s.graceCtx.Err() != nil {
|
||||
close(s.connHandlingDone)
|
||||
}
|
||||
}
|
||||
|
||||
// ServeQUICConn serves a single QUIC connection.
|
||||
func (s *Server) ServeQUICConn(conn *quic.Conn) error {
|
||||
s.mutex.Lock()
|
||||
if s.closed {
|
||||
s.mutex.Unlock()
|
||||
return http.ErrServerClosed
|
||||
}
|
||||
|
||||
s.init()
|
||||
s.mutex.Unlock()
|
||||
|
||||
s.connCount.Add(1)
|
||||
defer s.decreaseConnCount()
|
||||
|
||||
return s.handleConn(conn)
|
||||
}
|
||||
|
||||
// ServeListener serves an existing QUIC listener.
|
||||
// Make sure you use http3.ConfigureTLSConfig to configure a tls.Config
|
||||
// and use it to construct a http3-friendly QUIC listener.
|
||||
// Closing the server does not close the listener. It is the application's responsibility to close them.
|
||||
// ServeListener always returns a non-nil error. After Shutdown or Close, the returned error is http.ErrServerClosed.
|
||||
func (s *Server) ServeListener(ln QUICListener) error {
|
||||
s.mutex.Lock()
|
||||
if err := s.addListener(&ln, false); err != nil {
|
||||
s.mutex.Unlock()
|
||||
return err
|
||||
}
|
||||
s.mutex.Unlock()
|
||||
defer s.removeListener(&ln)
|
||||
|
||||
return s.serveListener(ln)
|
||||
}
|
||||
|
||||
func (s *Server) serveListener(ln QUICListener) error {
|
||||
for {
|
||||
conn, err := ln.Accept(s.graceCtx)
|
||||
// server closed
|
||||
if errors.Is(err, quic.ErrServerClosed) || s.graceCtx.Err() != nil {
|
||||
return http.ErrServerClosed
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.connCount.Add(1)
|
||||
go func() {
|
||||
defer s.decreaseConnCount()
|
||||
if err := s.handleConn(conn); err != nil {
|
||||
if s.Logger != nil {
|
||||
s.Logger.Debug("handling connection failed", "error", err)
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
var errServerWithoutTLSConfig = errors.New("use of http3.Server without TLSConfig")
|
||||
|
||||
func (s *Server) setupListenerForConn(tlsConf *tls.Config, conn net.PacketConn) (*QUICListener, error) {
|
||||
if tlsConf == nil {
|
||||
return nil, errServerWithoutTLSConfig
|
||||
}
|
||||
|
||||
baseConf := ConfigureTLSConfig(tlsConf)
|
||||
quicConf := s.QUICConfig
|
||||
if quicConf == nil {
|
||||
quicConf = &quic.Config{Allow0RTT: true}
|
||||
} else {
|
||||
quicConf = s.QUICConfig.Clone()
|
||||
}
|
||||
if s.EnableDatagrams {
|
||||
quicConf.EnableDatagrams = true
|
||||
}
|
||||
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
closed := s.closed
|
||||
if closed {
|
||||
return nil, http.ErrServerClosed
|
||||
}
|
||||
|
||||
var ln QUICListener
|
||||
var err error
|
||||
if conn == nil {
|
||||
addr := s.Addr
|
||||
if addr == "" {
|
||||
addr = ":https"
|
||||
}
|
||||
ln, err = quic.ListenAddrEarly(addr, baseConf, quicConf)
|
||||
} else {
|
||||
ln, err = quic.ListenEarly(conn, baseConf, quicConf)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.addListener(&ln, true); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &ln, nil
|
||||
}
|
||||
|
||||
func extractPort(addr string) (int, error) {
|
||||
_, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
portInt, err := net.LookupPort("tcp", portStr)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return portInt, nil
|
||||
}
|
||||
|
||||
func (s *Server) generateAltSvcHeader() {
|
||||
if len(s.listeners) == 0 {
|
||||
// Don't announce any ports since no one is listening for connections
|
||||
s.altSvcHeader = ""
|
||||
return
|
||||
}
|
||||
|
||||
// This code assumes that we will use protocol.SupportedVersions if no quic.Config is passed.
|
||||
|
||||
var altSvc []string
|
||||
addPort := func(port int) {
|
||||
altSvc = append(altSvc, fmt.Sprintf(`%s=":%d"; ma=2592000`, NextProtoH3, port))
|
||||
}
|
||||
|
||||
if s.Port != 0 {
|
||||
// if Port is specified, we must use it instead of the
|
||||
// listener addresses since there's a reason it's specified.
|
||||
addPort(s.Port)
|
||||
} else {
|
||||
// if we have some listeners assigned, try to find ports
|
||||
// which we can announce, otherwise nothing should be announced
|
||||
validPortsFound := false
|
||||
for _, info := range s.listeners {
|
||||
if info.port != 0 {
|
||||
addPort(info.port)
|
||||
validPortsFound = true
|
||||
}
|
||||
}
|
||||
if !validPortsFound {
|
||||
if port, err := extractPort(s.Addr); err == nil {
|
||||
addPort(port)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
s.altSvcHeader = strings.Join(altSvc, ",")
|
||||
}
|
||||
|
||||
func (s *Server) addListener(l *QUICListener, createdLocally bool) error {
|
||||
if s.closed {
|
||||
return http.ErrServerClosed
|
||||
}
|
||||
s.init()
|
||||
|
||||
laddr := (*l).Addr()
|
||||
if port, err := extractPort(laddr.String()); err == nil {
|
||||
s.listeners = append(s.listeners, listener{ln: l, port: port, createdLocally: createdLocally})
|
||||
} else {
|
||||
logger := s.Logger
|
||||
if logger == nil {
|
||||
logger = slog.Default()
|
||||
}
|
||||
logger.Error("Unable to extract port from listener, will not be announced using SetQUICHeaders", "local addr", laddr, "error", err)
|
||||
s.listeners = append(s.listeners, listener{ln: l, port: 0, createdLocally: createdLocally})
|
||||
}
|
||||
s.generateAltSvcHeader()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Server) removeListener(l *QUICListener) {
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
|
||||
s.listeners = slices.DeleteFunc(s.listeners, func(info listener) bool {
|
||||
return info.ln == l
|
||||
})
|
||||
s.generateAltSvcHeader()
|
||||
}
|
||||
|
||||
func (s *Server) NewRawServerConn(conn *quic.Conn) (*RawServerConn, error) {
|
||||
hconn, _, _, err := s.newRawServerConn(conn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return hconn, nil
|
||||
}
|
||||
|
||||
func (s *Server) newRawServerConn(conn *quic.Conn) (*RawServerConn, *quic.SendStream, qlogwriter.Recorder, error) {
|
||||
var qlogger qlogwriter.Recorder
|
||||
if qlogTrace := conn.QlogTrace(); qlogTrace != nil && qlogTrace.SupportsSchemas(qlog.EventSchema) {
|
||||
qlogger = qlogTrace.AddProducer()
|
||||
}
|
||||
connCtx := conn.Context()
|
||||
connCtx = context.WithValue(connCtx, ServerContextKey, s)
|
||||
connCtx = context.WithValue(connCtx, http.LocalAddrContextKey, conn.LocalAddr())
|
||||
connCtx = context.WithValue(connCtx, RemoteAddrContextKey, conn.RemoteAddr())
|
||||
if s.ConnContext != nil {
|
||||
connCtx = s.ConnContext(connCtx, conn)
|
||||
if connCtx == nil {
|
||||
panic("http3: ConnContext returned nil")
|
||||
}
|
||||
}
|
||||
hconn := newRawServerConn(
|
||||
conn,
|
||||
s.EnableDatagrams,
|
||||
s.IdleTimeout,
|
||||
qlogger,
|
||||
s.Logger,
|
||||
connCtx,
|
||||
s.Handler,
|
||||
s.maxHeaderBytes(),
|
||||
)
|
||||
|
||||
// open the control stream and send a SETTINGS frame, it's also used to send a GOAWAY frame later
|
||||
// when the server is gracefully closed
|
||||
ctrlStr, err := hconn.openControlStream(&settingsFrame{
|
||||
MaxFieldSectionSize: int64(s.maxHeaderBytes()),
|
||||
Datagram: s.EnableDatagrams,
|
||||
ExtendedConnect: true,
|
||||
Other: s.AdditionalSettings,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, nil, nil, fmt.Errorf("opening the control stream failed: %w", err)
|
||||
}
|
||||
return hconn, ctrlStr, qlogger, nil
|
||||
}
|
||||
|
||||
// handleConn handles the HTTP/3 exchange on a QUIC connection.
|
||||
// It blocks until all HTTP handlers for all streams have returned.
|
||||
func (s *Server) handleConn(conn *quic.Conn) error {
|
||||
hconn, ctrlStr, qlogger, err := s.newRawServerConn(conn)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Go(func() {
|
||||
for {
|
||||
str, err := conn.AcceptUniStream(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go hconn.HandleUnidirectionalStream(str)
|
||||
}
|
||||
})
|
||||
|
||||
var nextStreamID quic.StreamID
|
||||
var handleErr error
|
||||
var inGracefulShutdown bool
|
||||
// Process all requests immediately.
|
||||
// It's the client's responsibility to decide which requests are eligible for 0-RTT.
|
||||
ctx := s.graceCtx
|
||||
for {
|
||||
// The context used here is:
|
||||
// * before graceful shutdown: s.graceCtx
|
||||
// * after graceful shutdown: s.closeCtx
|
||||
// This allows us to keep accepting (and resetting) streams after graceful shutdown has started.
|
||||
str, err := conn.AcceptStream(ctx)
|
||||
if err != nil {
|
||||
// the underlying connection was closed (by either side)
|
||||
if conn.Context().Err() != nil {
|
||||
var appErr *quic.ApplicationError
|
||||
if !errors.As(err, &appErr) || appErr.ErrorCode != quic.ApplicationErrorCode(ErrCodeNoError) {
|
||||
handleErr = fmt.Errorf("accepting stream failed: %w", err)
|
||||
}
|
||||
break
|
||||
}
|
||||
// server (not gracefully) closed, close the connection immediately
|
||||
if s.closeCtx.Err() != nil {
|
||||
hconn.CloseWithError(quic.ApplicationErrorCode(ErrCodeNoError), "")
|
||||
handleErr = http.ErrServerClosed
|
||||
break
|
||||
}
|
||||
inGracefulShutdown = s.graceCtx.Err() != nil
|
||||
if !inGracefulShutdown {
|
||||
var appErr *quic.ApplicationError
|
||||
if !errors.As(err, &appErr) || appErr.ErrorCode != quic.ApplicationErrorCode(ErrCodeNoError) {
|
||||
handleErr = fmt.Errorf("accepting stream failed: %w", err)
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
// gracefully closed, send GOAWAY frame and wait for requests to complete or grace period to end
|
||||
// new requests will be rejected and shouldn't be sent
|
||||
if qlogger != nil {
|
||||
qlogger.RecordEvent(qlog.FrameCreated{
|
||||
StreamID: ctrlStr.StreamID(),
|
||||
Frame: qlog.Frame{Frame: qlog.GoAwayFrame{StreamID: nextStreamID}},
|
||||
})
|
||||
}
|
||||
// Send the GOAWAY frame in a separate Goroutine.
|
||||
// Sending might block if the peer didn't grant enough flow control credit.
|
||||
// Write is guaranteed to return once the connection is closed.
|
||||
wg.Go(func() {
|
||||
_, _ = ctrlStr.Write((&goAwayFrame{StreamID: nextStreamID}).Append(nil))
|
||||
})
|
||||
ctx = s.closeCtx
|
||||
continue
|
||||
}
|
||||
if inGracefulShutdown {
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeRequestRejected))
|
||||
str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestRejected))
|
||||
continue
|
||||
}
|
||||
|
||||
nextStreamID = str.StreamID() + 4
|
||||
wg.Go(func() {
|
||||
// HandleRequestStream will return once the request has been handled,
|
||||
// or the underlying connection is closed.
|
||||
hconn.HandleRequestStream(str)
|
||||
})
|
||||
}
|
||||
wg.Wait()
|
||||
return handleErr
|
||||
}
|
||||
|
||||
func (s *Server) maxHeaderBytes() int {
|
||||
if s.MaxHeaderBytes <= 0 {
|
||||
return http.DefaultMaxHeaderBytes
|
||||
}
|
||||
return s.MaxHeaderBytes
|
||||
}
|
||||
|
||||
// Close the server immediately, aborting requests and sending CONNECTION_CLOSE frames to connected clients.
|
||||
// Close in combination with ListenAndServe() (instead of Serve()) may race if it is called before a UDP socket is established.
|
||||
// It is the caller's responsibility to close any connection passed to ServeQUICConn.
|
||||
func (s *Server) Close() error {
|
||||
s.mutex.Lock()
|
||||
defer s.mutex.Unlock()
|
||||
|
||||
s.closed = true
|
||||
// server is never used
|
||||
if s.closeCtx == nil {
|
||||
return nil
|
||||
}
|
||||
s.closeCancel()
|
||||
|
||||
var err error
|
||||
for _, l := range s.listeners {
|
||||
if l.createdLocally {
|
||||
if cerr := (*l.ln).Close(); cerr != nil && err == nil {
|
||||
err = cerr
|
||||
}
|
||||
}
|
||||
}
|
||||
if s.connCount.Load() == 0 {
|
||||
return err
|
||||
}
|
||||
// wait for all connections to be closed
|
||||
<-s.connHandlingDone
|
||||
return err
|
||||
}
|
||||
|
||||
// Shutdown gracefully shuts down the server without interrupting any active connections.
|
||||
// The server sends a GOAWAY frame first, then or for all running requests to complete.
|
||||
// Shutdown in combination with ListenAndServe may race if it is called before a UDP socket is established.
|
||||
// It is recommended to use Serve instead.
|
||||
func (s *Server) Shutdown(ctx context.Context) error {
|
||||
s.mutex.Lock()
|
||||
s.closed = true
|
||||
// server was never used
|
||||
if s.closeCtx == nil {
|
||||
s.mutex.Unlock()
|
||||
return nil
|
||||
}
|
||||
s.graceCancel()
|
||||
|
||||
// close all listeners
|
||||
var closeErrs []error
|
||||
for _, l := range s.listeners {
|
||||
if l.createdLocally {
|
||||
if err := (*l.ln).Close(); err != nil {
|
||||
closeErrs = append(closeErrs, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
s.mutex.Unlock()
|
||||
if len(closeErrs) > 0 {
|
||||
return errors.Join(closeErrs...)
|
||||
}
|
||||
|
||||
if s.connCount.Load() == 0 {
|
||||
return s.Close()
|
||||
}
|
||||
select {
|
||||
case <-s.connHandlingDone: // all connections were closed
|
||||
// When receiving a GOAWAY frame, HTTP/3 clients are expected to close the connection
|
||||
// once all requests were successfully handled...
|
||||
return s.Close()
|
||||
case <-ctx.Done():
|
||||
// ... however, clients handling long-lived requests (and misbehaving clients),
|
||||
// might not do so before the context is cancelled.
|
||||
// In this case, we close the server, which closes all existing connections
|
||||
// (expect those passed to ServeQUICConn).
|
||||
_ = s.Close()
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
// ErrNoAltSvcPort is the error returned by SetQUICHeaders when no port was found
|
||||
// for Alt-Svc to announce. This can happen if listening on a PacketConn without a port
|
||||
// (UNIX socket, for example) and no port is specified in Server.Port or Server.Addr.
|
||||
var ErrNoAltSvcPort = errors.New("no port can be announced, specify it explicitly using Server.Port or Server.Addr")
|
||||
|
||||
// SetQUICHeaders can be used to set the proper headers that announce that this server supports HTTP/3.
|
||||
// The values set by default advertise all the ports the server is listening on, but can be
|
||||
// changed to a specific port by setting Server.Port before launching the server.
|
||||
// If no listener's Addr().String() returns an address with a valid port, Server.Addr will be used
|
||||
// to extract the port, if specified.
|
||||
// For example, a server launched using ListenAndServe on an address with port 443 would set:
|
||||
//
|
||||
// Alt-Svc: h3=":443"; ma=2592000
|
||||
func (s *Server) SetQUICHeaders(hdr http.Header) error {
|
||||
s.mutex.RLock()
|
||||
defer s.mutex.RUnlock()
|
||||
|
||||
if s.altSvcHeader == "" {
|
||||
return ErrNoAltSvcPort
|
||||
}
|
||||
// use the map directly to avoid constant canonicalization since the key is already canonicalized
|
||||
hdr["Alt-Svc"] = append(hdr["Alt-Svc"], s.altSvcHeader)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListenAndServeQUIC listens on the UDP network address addr and calls the
|
||||
// handler for HTTP/3 requests on incoming connections. http.DefaultServeMux is
|
||||
// used when handler is nil.
|
||||
func ListenAndServeQUIC(addr, certFile, keyFile string, handler http.Handler) error {
|
||||
server := &Server{
|
||||
Addr: addr,
|
||||
Handler: handler,
|
||||
}
|
||||
return server.ListenAndServeTLS(certFile, keyFile)
|
||||
}
|
||||
|
||||
// ListenAndServeTLS listens on the given network address for both TLS/TCP and QUIC
|
||||
// connections in parallel. It returns if one of the two returns an error.
|
||||
// http.DefaultServeMux is used when handler is nil.
|
||||
// The correct Alt-Svc headers for QUIC are set.
|
||||
func ListenAndServeTLS(addr, certFile, keyFile string, handler http.Handler) error {
|
||||
// Load certs
|
||||
var err error
|
||||
certs := make([]tls.Certificate, 1)
|
||||
certs[0], err = tls.LoadX509KeyPair(certFile, keyFile)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// We currently only use the cert-related stuff from tls.Config,
|
||||
// so we don't need to make a full copy.
|
||||
config := &tls.Config{
|
||||
Certificates: certs,
|
||||
}
|
||||
|
||||
if addr == "" {
|
||||
addr = ":https"
|
||||
}
|
||||
|
||||
// Open the listeners
|
||||
udpAddr, err := net.ResolveUDPAddr("udp", addr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
udpConn, err := net.ListenUDP("udp", udpAddr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer udpConn.Close()
|
||||
|
||||
if handler == nil {
|
||||
handler = http.DefaultServeMux
|
||||
}
|
||||
// Start the servers
|
||||
quicServer := &Server{
|
||||
TLSConfig: config,
|
||||
Handler: handler,
|
||||
}
|
||||
|
||||
hErr := make(chan error, 1)
|
||||
qErr := make(chan error, 1)
|
||||
go func() {
|
||||
hErr <- http.ListenAndServeTLS(addr, certFile, keyFile, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
quicServer.SetQUICHeaders(w.Header())
|
||||
handler.ServeHTTP(w, r)
|
||||
}))
|
||||
}()
|
||||
go func() {
|
||||
qErr <- quicServer.Serve(udpConn)
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-hErr:
|
||||
quicServer.Close()
|
||||
return err
|
||||
case err := <-qErr:
|
||||
// Cannot close the HTTP server or wait for requests to complete properly :/
|
||||
return err
|
||||
}
|
||||
}
|
||||
+261
@@ -0,0 +1,261 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"github.com/enetx/http"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/qpack"
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
)
|
||||
|
||||
// RawServerConn is an HTTP/3 server connection.
|
||||
// It can be used for advanced use cases where the application wants to manage the QUIC connection lifecycle.
|
||||
type RawServerConn struct {
|
||||
rawConn rawConn
|
||||
|
||||
idleTimeout time.Duration
|
||||
idleTimer *time.Timer
|
||||
|
||||
serverContext context.Context
|
||||
requestHandler http.Handler
|
||||
maxHeaderBytes int
|
||||
|
||||
decoder *qpack.Decoder
|
||||
|
||||
qlogger qlogwriter.Recorder
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
func newRawServerConn(
|
||||
conn *quic.Conn,
|
||||
enableDatagrams bool,
|
||||
idleTimeout time.Duration,
|
||||
qlogger qlogwriter.Recorder,
|
||||
logger *slog.Logger,
|
||||
serverContext context.Context,
|
||||
requestHandler http.Handler,
|
||||
maxHeaderBytes int,
|
||||
) *RawServerConn {
|
||||
c := &RawServerConn{
|
||||
idleTimeout: idleTimeout,
|
||||
serverContext: serverContext,
|
||||
requestHandler: requestHandler,
|
||||
maxHeaderBytes: maxHeaderBytes,
|
||||
decoder: qpack.NewDecoder(),
|
||||
qlogger: qlogger,
|
||||
logger: logger,
|
||||
}
|
||||
c.rawConn = *newRawConn(conn, enableDatagrams, c.onStreamsEmpty, nil, qlogger, logger)
|
||||
if idleTimeout > 0 {
|
||||
c.idleTimer = time.AfterFunc(idleTimeout, c.onIdleTimer)
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *RawServerConn) onStreamsEmpty() {
|
||||
if c.idleTimeout > 0 {
|
||||
c.idleTimer.Reset(c.idleTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *RawServerConn) onIdleTimer() {
|
||||
c.CloseWithError(quic.ApplicationErrorCode(ErrCodeNoError), "idle timeout")
|
||||
}
|
||||
|
||||
// CloseWithError closes the connection with the given error code and message.
|
||||
func (c *RawServerConn) CloseWithError(code quic.ApplicationErrorCode, msg string) error {
|
||||
if c.idleTimer != nil {
|
||||
c.idleTimer.Stop()
|
||||
}
|
||||
return c.rawConn.CloseWithError(code, msg)
|
||||
}
|
||||
|
||||
// HandleRequestStream handles an HTTP/3 request on a bidirectional request stream.
|
||||
// The stream can either be obtained by calling AcceptStream on the underlying QUIC connection,
|
||||
// or (internally) by using the server's stream accept loop.
|
||||
func (c *RawServerConn) HandleRequestStream(str *quic.Stream) {
|
||||
hstr := c.rawConn.TrackStream(str)
|
||||
c.handleRequestStream(hstr)
|
||||
}
|
||||
|
||||
func (c *RawServerConn) requestMaxHeaderBytes() int {
|
||||
if c.maxHeaderBytes <= 0 {
|
||||
return http.DefaultMaxHeaderBytes
|
||||
}
|
||||
return c.maxHeaderBytes
|
||||
}
|
||||
|
||||
func (c *RawServerConn) openControlStream(settings *settingsFrame) (*quic.SendStream, error) {
|
||||
return c.rawConn.openControlStream(settings)
|
||||
}
|
||||
|
||||
func (c *RawServerConn) handleRequestStream(str *stateTrackingStream) {
|
||||
if c.idleTimeout > 0 {
|
||||
// This only applies if the stream is the first active stream,
|
||||
// but it's ok to stop a stopped timer.
|
||||
c.idleTimer.Stop()
|
||||
}
|
||||
|
||||
conn := &c.rawConn
|
||||
qlogger := c.qlogger
|
||||
decoder := c.decoder
|
||||
connCtx := c.serverContext
|
||||
maxHeaderBytes := c.requestMaxHeaderBytes()
|
||||
|
||||
fp := &frameParser{closeConn: conn.CloseWithError, r: str, streamID: str.StreamID()}
|
||||
frame, err := fp.ParseNext(qlogger)
|
||||
if err != nil {
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeRequestIncomplete))
|
||||
str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestIncomplete))
|
||||
return
|
||||
}
|
||||
hf, ok := frame.(*headersFrame)
|
||||
if !ok {
|
||||
conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "expected first frame to be a HEADERS frame")
|
||||
return
|
||||
}
|
||||
if hf.Length > uint64(maxHeaderBytes) {
|
||||
maybeQlogInvalidHeadersFrame(qlogger, str.StreamID(), hf.Length)
|
||||
// stop the client from sending more data
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeExcessiveLoad))
|
||||
// send a 431 Response (Request Header Fields Too Large)
|
||||
c.rejectWithHeaderFieldsTooLarge(str)
|
||||
return
|
||||
}
|
||||
headerBlock := make([]byte, hf.Length)
|
||||
if _, err := io.ReadFull(str, headerBlock); err != nil {
|
||||
maybeQlogInvalidHeadersFrame(qlogger, str.StreamID(), hf.Length)
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeRequestIncomplete))
|
||||
str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestIncomplete))
|
||||
return
|
||||
}
|
||||
decodeFn := decoder.Decode(headerBlock)
|
||||
var hfs []qpack.HeaderField
|
||||
if qlogger != nil {
|
||||
hfs = make([]qpack.HeaderField, 0, 16)
|
||||
}
|
||||
req, err := requestFromHeaders(decodeFn, maxHeaderBytes, &hfs)
|
||||
if qlogger != nil {
|
||||
qlogParsedHeadersFrame(qlogger, str.StreamID(), hf, hfs)
|
||||
}
|
||||
if err != nil {
|
||||
if errors.Is(err, errHeaderTooLarge) {
|
||||
// stop the client from sending more data
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeExcessiveLoad))
|
||||
// send a 431 Response (Request Header Fields Too Large)
|
||||
c.rejectWithHeaderFieldsTooLarge(str)
|
||||
return
|
||||
}
|
||||
|
||||
errCode := ErrCodeMessageError
|
||||
var qpackErr *qpackError
|
||||
if errors.As(err, &qpackErr) {
|
||||
errCode = ErrCodeQPACKDecompressionFailed
|
||||
}
|
||||
str.CancelRead(quic.StreamErrorCode(errCode))
|
||||
str.CancelWrite(quic.StreamErrorCode(errCode))
|
||||
return
|
||||
}
|
||||
|
||||
connState := conn.ConnectionState().TLS
|
||||
req.TLS = &connState
|
||||
req.RemoteAddr = conn.RemoteAddr().String()
|
||||
|
||||
// Check that the client doesn't send more data in DATA frames than indicated by the Content-Length header (if set).
|
||||
// See section 4.1.2 of RFC 9114.
|
||||
contentLength := int64(-1)
|
||||
if _, ok := req.Header["Content-Length"]; ok && req.ContentLength >= 0 {
|
||||
contentLength = req.ContentLength
|
||||
}
|
||||
hstr := newStream(str, conn, nil, func(r io.Reader, hf *headersFrame) error {
|
||||
trailers, err := decodeTrailers(r, hf, maxHeaderBytes, decoder, qlogger, str.StreamID())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Trailer = trailers
|
||||
return nil
|
||||
}, qlogger)
|
||||
body := newRequestBody(hstr, contentLength, connCtx, conn.ReceivedSettings(), conn.Settings)
|
||||
req.Body = body
|
||||
|
||||
if c.logger != nil {
|
||||
c.logger.Debug("handling request", "method", req.Method, "host", req.Host, "uri", req.RequestURI)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(connCtx)
|
||||
req = req.WithContext(ctx)
|
||||
context.AfterFunc(str.Context(), cancel)
|
||||
|
||||
r := newResponseWriter(hstr, conn, req.Method == http.MethodHead, c.logger)
|
||||
handler := c.requestHandler
|
||||
if handler == nil {
|
||||
handler = http.DefaultServeMux
|
||||
}
|
||||
|
||||
// It's the client's responsibility to decide which requests are eligible for 0-RTT.
|
||||
var panicked bool
|
||||
func() {
|
||||
defer func() {
|
||||
if p := recover(); p != nil {
|
||||
panicked = true
|
||||
if p == http.ErrAbortHandler {
|
||||
return
|
||||
}
|
||||
// Copied from net/http/server.go
|
||||
const size = 64 << 10
|
||||
buf := make([]byte, size)
|
||||
buf = buf[:runtime.Stack(buf, false)]
|
||||
logger := c.logger
|
||||
if logger == nil {
|
||||
logger = slog.Default()
|
||||
}
|
||||
logger.Error("http3: panic serving", "arg", p, "trace", string(buf))
|
||||
}
|
||||
}()
|
||||
handler.ServeHTTP(r, req)
|
||||
}()
|
||||
|
||||
if r.wasStreamHijacked() {
|
||||
return
|
||||
}
|
||||
|
||||
// abort the stream when there is a panic
|
||||
if panicked {
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeInternalError))
|
||||
str.CancelWrite(quic.StreamErrorCode(ErrCodeInternalError))
|
||||
return
|
||||
}
|
||||
|
||||
// response not written to the client yet, set Content-Length
|
||||
if !r.headerWritten {
|
||||
if _, haveCL := r.header["Content-Length"]; !haveCL {
|
||||
r.header.Set("Content-Length", strconv.FormatInt(r.numWritten, 10))
|
||||
}
|
||||
}
|
||||
r.Flush()
|
||||
r.flushTrailers()
|
||||
|
||||
// If the EOF was read by the handler, CancelRead() is a no-op.
|
||||
str.CancelRead(quic.StreamErrorCode(ErrCodeNoError))
|
||||
str.Close()
|
||||
}
|
||||
|
||||
func (c *RawServerConn) rejectWithHeaderFieldsTooLarge(str *stateTrackingStream) {
|
||||
hstr := newStream(str, &c.rawConn, nil, nil, c.qlogger)
|
||||
defer hstr.Close()
|
||||
r := newResponseWriter(hstr, &c.rawConn, false, c.logger)
|
||||
r.WriteHeader(http.StatusRequestHeaderFieldsTooLarge)
|
||||
r.Flush()
|
||||
}
|
||||
|
||||
// HandleUnidirectionalStream handles an incoming unidirectional stream.
|
||||
func (c *RawServerConn) HandleUnidirectionalStream(str *quic.ReceiveStream) {
|
||||
c.rawConn.handleUnidirectionalStream(str, true)
|
||||
}
|
||||
+173
@@ -0,0 +1,173 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"sync"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
)
|
||||
|
||||
const streamDatagramQueueLen = 32
|
||||
|
||||
// stateTrackingStream is an implementation of quic.Stream that delegates
|
||||
// to an underlying stream
|
||||
// it takes care of proxying send and receive errors onto an implementation of
|
||||
// the errorSetter interface (intended to be occupied by a datagrammer)
|
||||
// it is also responsible for clearing the stream based on its ID from its
|
||||
// parent connection, this is done through the streamClearer interface when
|
||||
// both the send and receive sides are closed
|
||||
type stateTrackingStream struct {
|
||||
*quic.Stream
|
||||
|
||||
sendDatagram func([]byte) error
|
||||
hasData chan struct{}
|
||||
queue [][]byte // TODO: use a ring buffer
|
||||
|
||||
mx sync.Mutex
|
||||
sendErr error
|
||||
recvErr error
|
||||
|
||||
clearer streamClearer
|
||||
}
|
||||
|
||||
var _ datagramStream = &stateTrackingStream{}
|
||||
|
||||
type streamClearer interface {
|
||||
clearStream(quic.StreamID)
|
||||
}
|
||||
|
||||
func newStateTrackingStream(s *quic.Stream, clearer streamClearer, sendDatagram func([]byte) error) *stateTrackingStream {
|
||||
t := &stateTrackingStream{
|
||||
Stream: s,
|
||||
clearer: clearer,
|
||||
sendDatagram: sendDatagram,
|
||||
hasData: make(chan struct{}, 1),
|
||||
}
|
||||
|
||||
context.AfterFunc(s.Context(), func() {
|
||||
t.closeSend(context.Cause(s.Context()))
|
||||
})
|
||||
|
||||
return t
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) closeSend(e error) {
|
||||
s.mx.Lock()
|
||||
defer s.mx.Unlock()
|
||||
|
||||
// clear the stream the first time both the send
|
||||
// and receive are finished
|
||||
if s.sendErr == nil {
|
||||
if s.recvErr != nil {
|
||||
s.clearer.clearStream(s.StreamID())
|
||||
}
|
||||
s.sendErr = e
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) closeReceive(e error) {
|
||||
s.mx.Lock()
|
||||
defer s.mx.Unlock()
|
||||
|
||||
// clear the stream the first time both the send
|
||||
// and receive are finished
|
||||
if s.recvErr == nil {
|
||||
if s.sendErr != nil {
|
||||
s.clearer.clearStream(s.StreamID())
|
||||
}
|
||||
s.recvErr = e
|
||||
s.signalHasDatagram()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) Close() error {
|
||||
s.closeSend(errors.New("write on closed stream"))
|
||||
return s.Stream.Close()
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) CancelWrite(e quic.StreamErrorCode) {
|
||||
s.closeSend(&quic.StreamError{StreamID: s.StreamID(), ErrorCode: e})
|
||||
s.Stream.CancelWrite(e)
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) Write(b []byte) (int, error) {
|
||||
n, err := s.Stream.Write(b)
|
||||
if err != nil && !errors.Is(err, os.ErrDeadlineExceeded) {
|
||||
s.closeSend(err)
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) CancelRead(e quic.StreamErrorCode) {
|
||||
s.closeReceive(&quic.StreamError{StreamID: s.StreamID(), ErrorCode: e})
|
||||
s.Stream.CancelRead(e)
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) Read(b []byte) (int, error) {
|
||||
n, err := s.Stream.Read(b)
|
||||
if err != nil && !errors.Is(err, os.ErrDeadlineExceeded) {
|
||||
s.closeReceive(err)
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) SendDatagram(b []byte) error {
|
||||
s.mx.Lock()
|
||||
sendErr := s.sendErr
|
||||
s.mx.Unlock()
|
||||
if sendErr != nil {
|
||||
return sendErr
|
||||
}
|
||||
|
||||
return s.sendDatagram(b)
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) signalHasDatagram() {
|
||||
select {
|
||||
case s.hasData <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) enqueueDatagram(data []byte) {
|
||||
s.mx.Lock()
|
||||
defer s.mx.Unlock()
|
||||
|
||||
if s.recvErr != nil {
|
||||
return
|
||||
}
|
||||
if len(s.queue) >= streamDatagramQueueLen {
|
||||
return
|
||||
}
|
||||
s.queue = append(s.queue, data)
|
||||
s.signalHasDatagram()
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) ReceiveDatagram(ctx context.Context) ([]byte, error) {
|
||||
start:
|
||||
s.mx.Lock()
|
||||
if len(s.queue) > 0 {
|
||||
data := s.queue[0]
|
||||
s.queue = s.queue[1:]
|
||||
s.mx.Unlock()
|
||||
return data, nil
|
||||
}
|
||||
if receiveErr := s.recvErr; receiveErr != nil {
|
||||
s.mx.Unlock()
|
||||
return nil, receiveErr
|
||||
}
|
||||
s.mx.Unlock()
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, context.Cause(ctx)
|
||||
case <-s.hasData:
|
||||
}
|
||||
goto start
|
||||
}
|
||||
|
||||
func (s *stateTrackingStream) QUICStream() *quic.Stream {
|
||||
return s.Stream
|
||||
}
|
||||
+406
@@ -0,0 +1,406 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"github.com/enetx/http"
|
||||
"github.com/enetx/http/httptrace"
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/http3/qlog"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
|
||||
"github.com/quic-go/qpack"
|
||||
)
|
||||
|
||||
type datagramStream interface {
|
||||
io.ReadWriteCloser
|
||||
CancelRead(quic.StreamErrorCode)
|
||||
CancelWrite(quic.StreamErrorCode)
|
||||
StreamID() quic.StreamID
|
||||
Context() context.Context
|
||||
SetDeadline(time.Time) error
|
||||
SetReadDeadline(time.Time) error
|
||||
SetWriteDeadline(time.Time) error
|
||||
SendDatagram(b []byte) error
|
||||
ReceiveDatagram(ctx context.Context) ([]byte, error)
|
||||
|
||||
QUICStream() *quic.Stream
|
||||
}
|
||||
|
||||
// A Stream is an HTTP/3 stream.
|
||||
//
|
||||
// When writing to and reading from the stream, data is framed in HTTP/3 DATA frames.
|
||||
type Stream struct {
|
||||
datagramStream
|
||||
conn *rawConn
|
||||
frameParser *frameParser
|
||||
|
||||
buf []byte // used as a temporary buffer when writing the HTTP/3 frame headers
|
||||
|
||||
bytesRemainingInFrame uint64
|
||||
|
||||
qlogger qlogwriter.Recorder
|
||||
|
||||
parseTrailer func(io.Reader, *headersFrame) error
|
||||
parsedTrailer bool
|
||||
}
|
||||
|
||||
func newStream(
|
||||
str datagramStream,
|
||||
conn *rawConn,
|
||||
trace *httptrace.ClientTrace,
|
||||
parseTrailer func(io.Reader, *headersFrame) error,
|
||||
qlogger qlogwriter.Recorder,
|
||||
) *Stream {
|
||||
return &Stream{
|
||||
datagramStream: str,
|
||||
conn: conn,
|
||||
buf: make([]byte, 16),
|
||||
qlogger: qlogger,
|
||||
parseTrailer: parseTrailer,
|
||||
frameParser: &frameParser{
|
||||
r: &tracingReader{Reader: str, trace: trace},
|
||||
streamID: str.StreamID(),
|
||||
closeConn: conn.CloseWithError,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Stream) Read(b []byte) (int, error) {
|
||||
if s.bytesRemainingInFrame == 0 {
|
||||
parseLoop:
|
||||
for {
|
||||
frame, err := s.frameParser.ParseNext(s.qlogger)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
switch f := frame.(type) {
|
||||
case *dataFrame:
|
||||
if s.parsedTrailer {
|
||||
return 0, errors.New("DATA frame received after trailers")
|
||||
}
|
||||
s.bytesRemainingInFrame = f.Length
|
||||
break parseLoop
|
||||
case *headersFrame:
|
||||
if s.parsedTrailer {
|
||||
maybeQlogInvalidHeadersFrame(s.qlogger, s.StreamID(), f.Length)
|
||||
return 0, errors.New("additional HEADERS frame received after trailers")
|
||||
}
|
||||
s.parsedTrailer = true
|
||||
return 0, s.parseTrailer(s.datagramStream, f)
|
||||
default:
|
||||
s.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "")
|
||||
// parseNextFrame skips over unknown frame types
|
||||
// Therefore, this condition is only entered when we parsed another known frame type.
|
||||
return 0, fmt.Errorf("peer sent an unexpected frame: %T", f)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var n int
|
||||
var err error
|
||||
if s.bytesRemainingInFrame < uint64(len(b)) {
|
||||
n, err = s.datagramStream.Read(b[:s.bytesRemainingInFrame])
|
||||
} else {
|
||||
n, err = s.datagramStream.Read(b)
|
||||
}
|
||||
s.bytesRemainingInFrame -= uint64(n)
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (s *Stream) hasMoreData() bool {
|
||||
return s.bytesRemainingInFrame > 0
|
||||
}
|
||||
|
||||
func (s *Stream) Write(b []byte) (int, error) {
|
||||
s.buf = s.buf[:0]
|
||||
s.buf = (&dataFrame{Length: uint64(len(b))}).Append(s.buf)
|
||||
if s.qlogger != nil {
|
||||
s.qlogger.RecordEvent(qlog.FrameCreated{
|
||||
StreamID: s.StreamID(),
|
||||
Raw: qlog.RawInfo{
|
||||
Length: len(s.buf) + len(b),
|
||||
PayloadLength: len(b),
|
||||
},
|
||||
Frame: qlog.Frame{Frame: qlog.DataFrame{}},
|
||||
})
|
||||
}
|
||||
if _, err := s.datagramStream.Write(s.buf); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return s.datagramStream.Write(b)
|
||||
}
|
||||
|
||||
func (s *Stream) writeUnframed(b []byte) (int, error) {
|
||||
return s.datagramStream.Write(b)
|
||||
}
|
||||
|
||||
func (s *Stream) StreamID() quic.StreamID {
|
||||
return s.datagramStream.StreamID()
|
||||
}
|
||||
|
||||
func (s *Stream) SendDatagram(b []byte) error {
|
||||
// TODO: reject if datagrams are not negotiated (yet)
|
||||
return s.datagramStream.SendDatagram(b)
|
||||
}
|
||||
|
||||
func (s *Stream) ReceiveDatagram(ctx context.Context) ([]byte, error) {
|
||||
// TODO: reject if datagrams are not negotiated (yet)
|
||||
return s.datagramStream.ReceiveDatagram(ctx)
|
||||
}
|
||||
|
||||
// A RequestStream is a low-level abstraction representing an HTTP/3 request stream.
|
||||
// It decouples sending of the HTTP request from reading the HTTP response, allowing
|
||||
// the application to optimistically use the stream (and, for example, send datagrams)
|
||||
// before receiving the response.
|
||||
//
|
||||
// This is only needed for advanced use case, e.g. WebTransport and the various
|
||||
// MASQUE proxying protocols.
|
||||
type RequestStream struct {
|
||||
str *Stream
|
||||
|
||||
responseBody io.ReadCloser // set by ReadResponse
|
||||
|
||||
decoder *qpack.Decoder
|
||||
requestWriter *requestWriter
|
||||
maxHeaderBytes int
|
||||
reqDone chan<- struct{}
|
||||
disableCompression bool
|
||||
response *http.Response
|
||||
|
||||
sentRequest bool
|
||||
requestedGzip bool
|
||||
isConnect bool
|
||||
}
|
||||
|
||||
func newRequestStream(
|
||||
str *Stream,
|
||||
requestWriter *requestWriter,
|
||||
reqDone chan<- struct{},
|
||||
decoder *qpack.Decoder,
|
||||
disableCompression bool,
|
||||
maxHeaderBytes int,
|
||||
rsp *http.Response,
|
||||
) *RequestStream {
|
||||
return &RequestStream{
|
||||
str: str,
|
||||
requestWriter: requestWriter,
|
||||
reqDone: reqDone,
|
||||
decoder: decoder,
|
||||
disableCompression: disableCompression,
|
||||
maxHeaderBytes: maxHeaderBytes,
|
||||
response: rsp,
|
||||
}
|
||||
}
|
||||
|
||||
// Read reads data from the underlying stream.
|
||||
//
|
||||
// It can only be used after the request has been sent (using SendRequestHeader)
|
||||
// and the response has been consumed (using ReadResponse).
|
||||
func (s *RequestStream) Read(b []byte) (int, error) {
|
||||
if s.responseBody == nil {
|
||||
return 0, errors.New("http3: invalid use of RequestStream.Read before ReadResponse")
|
||||
}
|
||||
return s.responseBody.Read(b)
|
||||
}
|
||||
|
||||
// StreamID returns the QUIC stream ID of the underlying QUIC stream.
|
||||
func (s *RequestStream) StreamID() quic.StreamID {
|
||||
return s.str.StreamID()
|
||||
}
|
||||
|
||||
// Write writes data to the stream.
|
||||
//
|
||||
// It can only be used after the request has been sent (using SendRequestHeader).
|
||||
func (s *RequestStream) Write(b []byte) (int, error) {
|
||||
if !s.sentRequest {
|
||||
return 0, errors.New("http3: invalid use of RequestStream.Write before SendRequestHeader")
|
||||
}
|
||||
return s.str.Write(b)
|
||||
}
|
||||
|
||||
// Close closes the send-direction of the stream.
|
||||
// It does not close the receive-direction of the stream.
|
||||
func (s *RequestStream) Close() error {
|
||||
return s.str.Close()
|
||||
}
|
||||
|
||||
// CancelRead aborts receiving on this stream.
|
||||
// See [quic.Stream.CancelRead] for more details.
|
||||
func (s *RequestStream) CancelRead(errorCode quic.StreamErrorCode) {
|
||||
s.str.CancelRead(errorCode)
|
||||
}
|
||||
|
||||
// CancelWrite aborts sending on this stream.
|
||||
// See [quic.Stream.CancelWrite] for more details.
|
||||
func (s *RequestStream) CancelWrite(errorCode quic.StreamErrorCode) {
|
||||
s.str.CancelWrite(errorCode)
|
||||
}
|
||||
|
||||
// Context returns a context derived from the underlying QUIC stream's context.
|
||||
// See [quic.Stream.Context] for more details.
|
||||
func (s *RequestStream) Context() context.Context {
|
||||
return s.str.Context()
|
||||
}
|
||||
|
||||
// SetReadDeadline sets the deadline for Read calls.
|
||||
func (s *RequestStream) SetReadDeadline(t time.Time) error {
|
||||
return s.str.SetReadDeadline(t)
|
||||
}
|
||||
|
||||
// SetWriteDeadline sets the deadline for Write calls.
|
||||
func (s *RequestStream) SetWriteDeadline(t time.Time) error {
|
||||
return s.str.SetWriteDeadline(t)
|
||||
}
|
||||
|
||||
// SetDeadline sets the read and write deadlines associated with the stream.
|
||||
// It is equivalent to calling both SetReadDeadline and SetWriteDeadline.
|
||||
func (s *RequestStream) SetDeadline(t time.Time) error {
|
||||
return s.str.SetDeadline(t)
|
||||
}
|
||||
|
||||
// SendDatagrams send a new HTTP Datagram (RFC 9297).
|
||||
//
|
||||
// It is only possible to send datagrams if the server enabled support for this extension.
|
||||
// It is recommended (though not required) to send the request before calling this method,
|
||||
// as the server might drop datagrams which it can't associate with an existing request.
|
||||
func (s *RequestStream) SendDatagram(b []byte) error {
|
||||
return s.str.SendDatagram(b)
|
||||
}
|
||||
|
||||
// ReceiveDatagram receives HTTP Datagrams (RFC 9297).
|
||||
//
|
||||
// It is only possible if support for HTTP Datagrams was enabled, using the EnableDatagram
|
||||
// option on the [Transport].
|
||||
func (s *RequestStream) ReceiveDatagram(ctx context.Context) ([]byte, error) {
|
||||
return s.str.ReceiveDatagram(ctx)
|
||||
}
|
||||
|
||||
// SendRequestHeader sends the HTTP request.
|
||||
//
|
||||
// It can only used for requests that don't have a request body.
|
||||
// It is invalid to call it more than once.
|
||||
// It is invalid to call it after Write has been called.
|
||||
func (s *RequestStream) SendRequestHeader(req *http.Request) error {
|
||||
if req.Body != nil && req.Body != http.NoBody {
|
||||
return errors.New("http3: invalid use of RequestStream.SendRequestHeader with a request that has a request body")
|
||||
}
|
||||
return s.sendRequestHeader(req)
|
||||
}
|
||||
|
||||
func (s *RequestStream) sendRequestHeader(req *http.Request) error {
|
||||
if s.sentRequest {
|
||||
return errors.New("http3: invalid duplicate use of RequestStream.SendRequestHeader")
|
||||
}
|
||||
if !s.disableCompression && req.Method != http.MethodHead &&
|
||||
req.Header.Get("Accept-Encoding") == "" && req.Header.Get("Range") == "" {
|
||||
s.requestedGzip = true
|
||||
}
|
||||
s.isConnect = req.Method == http.MethodConnect
|
||||
s.sentRequest = true
|
||||
return s.requestWriter.WriteRequestHeader(s.str.datagramStream, req, s.requestedGzip, s.str.StreamID(), s.str.qlogger)
|
||||
}
|
||||
|
||||
// sendRequestTrailer sends request trailers to the stream.
|
||||
// It should be called after the request body has been fully written.
|
||||
func (s *RequestStream) sendRequestTrailer(req *http.Request) error {
|
||||
return s.requestWriter.WriteRequestTrailer(s.str.datagramStream, req, s.str.StreamID(), s.str.qlogger)
|
||||
}
|
||||
|
||||
// ReadResponse reads the HTTP response from the stream.
|
||||
//
|
||||
// It must be called after sending the request (using SendRequestHeader).
|
||||
// It is invalid to call it more than once.
|
||||
// It doesn't set Response.Request and Response.TLS.
|
||||
// It is invalid to call it after Read has been called.
|
||||
func (s *RequestStream) ReadResponse() (*http.Response, error) {
|
||||
if !s.sentRequest {
|
||||
return nil, errors.New("http3: invalid use of RequestStream.ReadResponse before SendRequestHeader")
|
||||
}
|
||||
frame, err := s.str.frameParser.ParseNext(s.str.qlogger)
|
||||
if err != nil {
|
||||
s.str.CancelRead(quic.StreamErrorCode(ErrCodeFrameError))
|
||||
s.str.CancelWrite(quic.StreamErrorCode(ErrCodeFrameError))
|
||||
return nil, fmt.Errorf("http3: parsing frame failed: %w", err)
|
||||
}
|
||||
hf, ok := frame.(*headersFrame)
|
||||
if !ok {
|
||||
s.str.conn.CloseWithError(quic.ApplicationErrorCode(ErrCodeFrameUnexpected), "expected first frame to be a HEADERS frame")
|
||||
return nil, errors.New("http3: expected first frame to be a HEADERS frame")
|
||||
}
|
||||
if hf.Length > uint64(s.maxHeaderBytes) {
|
||||
maybeQlogInvalidHeadersFrame(s.str.qlogger, s.str.StreamID(), hf.Length)
|
||||
s.str.CancelRead(quic.StreamErrorCode(ErrCodeFrameError))
|
||||
s.str.CancelWrite(quic.StreamErrorCode(ErrCodeFrameError))
|
||||
return nil, fmt.Errorf("http3: HEADERS frame too large: %d bytes (max: %d)", hf.Length, s.maxHeaderBytes)
|
||||
}
|
||||
headerBlock := make([]byte, hf.Length)
|
||||
if _, err := io.ReadFull(s.str.datagramStream, headerBlock); err != nil {
|
||||
maybeQlogInvalidHeadersFrame(s.str.qlogger, s.str.StreamID(), hf.Length)
|
||||
s.str.CancelRead(quic.StreamErrorCode(ErrCodeRequestIncomplete))
|
||||
s.str.CancelWrite(quic.StreamErrorCode(ErrCodeRequestIncomplete))
|
||||
return nil, fmt.Errorf("http3: failed to read response headers: %w", err)
|
||||
}
|
||||
decodeFn := s.decoder.Decode(headerBlock)
|
||||
var hfs []qpack.HeaderField
|
||||
if s.str.qlogger != nil {
|
||||
hfs = make([]qpack.HeaderField, 0, 16)
|
||||
}
|
||||
res := s.response
|
||||
err = updateResponseFromHeaders(res, decodeFn, s.maxHeaderBytes, &hfs)
|
||||
if s.str.qlogger != nil {
|
||||
qlogParsedHeadersFrame(s.str.qlogger, s.str.StreamID(), hf, hfs)
|
||||
}
|
||||
if err != nil {
|
||||
errCode := ErrCodeMessageError
|
||||
var qpackErr *qpackError
|
||||
if errors.As(err, &qpackErr) {
|
||||
errCode = ErrCodeQPACKDecompressionFailed
|
||||
}
|
||||
s.str.CancelRead(quic.StreamErrorCode(errCode))
|
||||
s.str.CancelWrite(quic.StreamErrorCode(errCode))
|
||||
return nil, fmt.Errorf("http3: invalid response: %w", err)
|
||||
}
|
||||
|
||||
// Check that the server doesn't send more data in DATA frames than indicated by the Content-Length header (if set).
|
||||
// See section 4.1.2 of RFC 9114.
|
||||
respBody := newResponseBody(s.str, res.ContentLength, s.reqDone)
|
||||
|
||||
// Rules for when to set Content-Length are defined in https://tools.ietf.org/html/rfc7230#section-3.3.2.
|
||||
isInformational := res.StatusCode >= 100 && res.StatusCode < 200
|
||||
isNoContent := res.StatusCode == http.StatusNoContent
|
||||
isSuccessfulConnect := s.isConnect && res.StatusCode >= 200 && res.StatusCode < 300
|
||||
if (isInformational || isNoContent || isSuccessfulConnect) && res.ContentLength == -1 {
|
||||
res.ContentLength = 0
|
||||
}
|
||||
if s.requestedGzip && res.Header.Get("Content-Encoding") == "gzip" {
|
||||
res.Header.Del("Content-Encoding")
|
||||
res.Header.Del("Content-Length")
|
||||
res.ContentLength = -1
|
||||
s.responseBody = newGzipReader(respBody)
|
||||
res.Uncompressed = true
|
||||
} else {
|
||||
s.responseBody = respBody
|
||||
}
|
||||
res.Body = s.responseBody
|
||||
return res, nil
|
||||
}
|
||||
|
||||
type tracingReader struct {
|
||||
io.Reader
|
||||
readFirst bool
|
||||
trace *httptrace.ClientTrace
|
||||
}
|
||||
|
||||
func (r *tracingReader) Read(b []byte) (int, error) {
|
||||
n, err := r.Reader.Read(b)
|
||||
if n > 0 && !r.readFirst {
|
||||
traceGotFirstResponseByte(r.trace)
|
||||
r.readFirst = true
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
+105
@@ -0,0 +1,105 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"net"
|
||||
"github.com/enetx/http/httptrace"
|
||||
"net/textproto"
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
)
|
||||
|
||||
func traceGetConn(trace *httptrace.ClientTrace, hostPort string) {
|
||||
if trace != nil && trace.GetConn != nil {
|
||||
trace.GetConn(hostPort)
|
||||
}
|
||||
}
|
||||
|
||||
// fakeConn is a wrapper for quic.EarlyConnection
|
||||
// because the quic connection does not implement net.Conn.
|
||||
type fakeConn struct {
|
||||
conn *quic.Conn
|
||||
}
|
||||
|
||||
func (c *fakeConn) Close() error { panic("connection operation prohibited") }
|
||||
func (c *fakeConn) Read(p []byte) (int, error) { panic("connection operation prohibited") }
|
||||
func (c *fakeConn) Write(p []byte) (int, error) { panic("connection operation prohibited") }
|
||||
func (c *fakeConn) SetDeadline(t time.Time) error { panic("connection operation prohibited") }
|
||||
func (c *fakeConn) SetReadDeadline(t time.Time) error { panic("connection operation prohibited") }
|
||||
func (c *fakeConn) SetWriteDeadline(t time.Time) error { panic("connection operation prohibited") }
|
||||
func (c *fakeConn) RemoteAddr() net.Addr { return c.conn.RemoteAddr() }
|
||||
func (c *fakeConn) LocalAddr() net.Addr { return c.conn.LocalAddr() }
|
||||
|
||||
func traceGotConn(trace *httptrace.ClientTrace, conn *quic.Conn, reused bool) {
|
||||
if trace != nil && trace.GotConn != nil {
|
||||
trace.GotConn(httptrace.GotConnInfo{
|
||||
Conn: &fakeConn{conn: conn},
|
||||
Reused: reused,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func traceGotFirstResponseByte(trace *httptrace.ClientTrace) {
|
||||
if trace != nil && trace.GotFirstResponseByte != nil {
|
||||
trace.GotFirstResponseByte()
|
||||
}
|
||||
}
|
||||
|
||||
func traceGot1xxResponse(trace *httptrace.ClientTrace, code int, header textproto.MIMEHeader) {
|
||||
if trace != nil && trace.Got1xxResponse != nil {
|
||||
trace.Got1xxResponse(code, header)
|
||||
}
|
||||
}
|
||||
|
||||
func traceGot100Continue(trace *httptrace.ClientTrace) {
|
||||
if trace != nil && trace.Got100Continue != nil {
|
||||
trace.Got100Continue()
|
||||
}
|
||||
}
|
||||
|
||||
func traceHasWroteHeaderField(trace *httptrace.ClientTrace) bool {
|
||||
return trace != nil && trace.WroteHeaderField != nil
|
||||
}
|
||||
|
||||
func traceWroteHeaderField(trace *httptrace.ClientTrace, k, v string) {
|
||||
if trace != nil && trace.WroteHeaderField != nil {
|
||||
trace.WroteHeaderField(k, []string{v})
|
||||
}
|
||||
}
|
||||
|
||||
func traceWroteHeaders(trace *httptrace.ClientTrace) {
|
||||
if trace != nil && trace.WroteHeaders != nil {
|
||||
trace.WroteHeaders()
|
||||
}
|
||||
}
|
||||
|
||||
func traceWroteRequest(trace *httptrace.ClientTrace, err error) {
|
||||
if trace != nil && trace.WroteRequest != nil {
|
||||
trace.WroteRequest(httptrace.WroteRequestInfo{Err: err})
|
||||
}
|
||||
}
|
||||
|
||||
func traceConnectStart(trace *httptrace.ClientTrace, network, addr string) {
|
||||
if trace != nil && trace.ConnectStart != nil {
|
||||
trace.ConnectStart(network, addr)
|
||||
}
|
||||
}
|
||||
|
||||
func traceConnectDone(trace *httptrace.ClientTrace, network, addr string, err error) {
|
||||
if trace != nil && trace.ConnectDone != nil {
|
||||
trace.ConnectDone(network, addr, err)
|
||||
}
|
||||
}
|
||||
|
||||
func traceTLSHandshakeStart(trace *httptrace.ClientTrace) {
|
||||
if trace != nil && trace.TLSHandshakeStart != nil {
|
||||
trace.TLSHandshakeStart()
|
||||
}
|
||||
}
|
||||
|
||||
func traceTLSHandshakeDone(trace *httptrace.ClientTrace, state tls.ConnectionState, err error) {
|
||||
if trace != nil && trace.TLSHandshakeDone != nil {
|
||||
trace.TLSHandshakeDone(state, err)
|
||||
}
|
||||
}
|
||||
+545
@@ -0,0 +1,545 @@
|
||||
package http3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/http"
|
||||
"github.com/enetx/http/httptrace"
|
||||
|
||||
"golang.org/x/net/http/httpguts"
|
||||
|
||||
"github.com/enetx/http3/httpcommon"
|
||||
"github.com/quic-go/quic-go"
|
||||
)
|
||||
|
||||
// Settings are HTTP/3 settings that apply to the underlying connection.
|
||||
type Settings struct {
|
||||
// Support for HTTP/3 datagrams (RFC 9297)
|
||||
EnableDatagrams bool
|
||||
// Extended CONNECT, RFC 9220
|
||||
EnableExtendedConnect bool
|
||||
// Other settings, defined by the application
|
||||
Other g.MapOrd[uint64, uint64]
|
||||
}
|
||||
|
||||
// RoundTripOpt are options for the Transport.RoundTripOpt method.
|
||||
type RoundTripOpt struct {
|
||||
// OnlyCachedConn controls whether the Transport may create a new QUIC connection.
|
||||
// If set true and no cached connection is available, RoundTripOpt will return ErrNoCachedConn.
|
||||
OnlyCachedConn bool
|
||||
}
|
||||
|
||||
type clientConn interface {
|
||||
OpenRequestStream(context.Context) (*RequestStream, error)
|
||||
RoundTrip(*http.Request) (*http.Response, error)
|
||||
handleUnidirectionalStream(*quic.ReceiveStream)
|
||||
}
|
||||
|
||||
type roundTripperWithCount struct {
|
||||
cancel context.CancelFunc
|
||||
dialing chan struct{} // closed as soon as quic.Dial(Early) returned
|
||||
dialErr error
|
||||
conn *quic.Conn
|
||||
clientConn clientConn
|
||||
|
||||
useCount atomic.Int64
|
||||
}
|
||||
|
||||
func (r *roundTripperWithCount) Close() error {
|
||||
r.cancel()
|
||||
<-r.dialing
|
||||
if r.conn != nil {
|
||||
return r.conn.CloseWithError(0, "")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Transport implements the http.RoundTripper interface
|
||||
type Transport struct {
|
||||
// TLSClientConfig specifies the TLS configuration to use with
|
||||
// tls.Client. If nil, the default configuration is used.
|
||||
TLSClientConfig *tls.Config
|
||||
|
||||
// QUICConfig is the quic.Config used for dialing new connections.
|
||||
// If nil, reasonable default values will be used.
|
||||
QUICConfig *quic.Config
|
||||
|
||||
// Dial specifies an optional dial function for creating QUIC
|
||||
// connections for requests.
|
||||
// If Dial is nil, a UDPConn will be created at the first request
|
||||
// and will be reused for subsequent connections to other servers.
|
||||
Dial func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error)
|
||||
|
||||
// Enable support for HTTP/3 datagrams (RFC 9297).
|
||||
// If a QUICConfig is set, datagram support also needs to be enabled on the QUIC layer by setting EnableDatagrams.
|
||||
EnableDatagrams bool
|
||||
|
||||
// Additional HTTP/3 settings.
|
||||
// It is invalid to specify any settings defined by RFC 9114 (HTTP/3) and RFC 9297 (HTTP Datagrams).
|
||||
AdditionalSettings g.MapOrd[uint64, uint64]
|
||||
|
||||
// MaxResponseHeaderBytes specifies a limit on how many response bytes are
|
||||
// allowed in the server's response header.
|
||||
// Zero means to use a default limit.
|
||||
MaxResponseHeaderBytes int
|
||||
|
||||
// DisableCompression, if true, prevents the Transport from requesting compression with an
|
||||
// "Accept-Encoding: gzip" request header when the Request contains no existing Accept-Encoding value.
|
||||
// If the Transport requests gzip on its own and gets a gzipped response, it's transparently
|
||||
// decoded in the Response.Body.
|
||||
// However, if the user explicitly requested gzip it is not automatically uncompressed.
|
||||
DisableCompression bool
|
||||
|
||||
Logger *slog.Logger
|
||||
|
||||
mutex sync.Mutex
|
||||
|
||||
initOnce sync.Once
|
||||
initErr error
|
||||
|
||||
newClientConn func(*quic.Conn) clientConn
|
||||
|
||||
clients map[string]*roundTripperWithCount
|
||||
transport *quic.Transport
|
||||
closed bool
|
||||
}
|
||||
|
||||
var (
|
||||
_ http.RoundTripper = &Transport{}
|
||||
_ io.Closer = &Transport{}
|
||||
)
|
||||
|
||||
var (
|
||||
// ErrNoCachedConn is returned when Transport.OnlyCachedConn is set
|
||||
ErrNoCachedConn = errors.New("http3: no cached connection was available")
|
||||
// ErrTransportClosed is returned when attempting to use a closed Transport
|
||||
ErrTransportClosed = errors.New("http3: transport is closed")
|
||||
)
|
||||
|
||||
func (t *Transport) init() error {
|
||||
if t.newClientConn == nil {
|
||||
t.newClientConn = func(conn *quic.Conn) clientConn {
|
||||
return newClientConn(
|
||||
conn,
|
||||
t.EnableDatagrams,
|
||||
t.AdditionalSettings,
|
||||
t.MaxResponseHeaderBytes,
|
||||
t.DisableCompression,
|
||||
t.Logger,
|
||||
)
|
||||
}
|
||||
}
|
||||
if t.QUICConfig == nil {
|
||||
t.QUICConfig = defaultQuicConfig.Clone()
|
||||
t.QUICConfig.EnableDatagrams = t.EnableDatagrams
|
||||
}
|
||||
if t.EnableDatagrams && !t.QUICConfig.EnableDatagrams {
|
||||
return errors.New("HTTP Datagrams enabled, but QUIC Datagrams disabled")
|
||||
}
|
||||
if len(t.QUICConfig.Versions) == 0 {
|
||||
t.QUICConfig = t.QUICConfig.Clone()
|
||||
t.QUICConfig.Versions = []quic.Version{quic.SupportedVersions()[0]}
|
||||
}
|
||||
if len(t.QUICConfig.Versions) != 1 {
|
||||
return errors.New("can only use a single QUIC version for dialing a HTTP/3 connection")
|
||||
}
|
||||
if t.QUICConfig.MaxIncomingStreams == 0 {
|
||||
t.QUICConfig.MaxIncomingStreams = -1 // don't allow any bidirectional streams
|
||||
}
|
||||
if t.Dial == nil {
|
||||
udpConn, err := net.ListenUDP("udp", nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
t.transport = &quic.Transport{Conn: udpConn}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RoundTripOpt is like RoundTrip, but takes options.
|
||||
func (t *Transport) RoundTripOpt(req *http.Request, opt RoundTripOpt) (*http.Response, error) {
|
||||
rsp, err := t.roundTripOpt(req, opt)
|
||||
if err != nil {
|
||||
if req.Body != nil {
|
||||
req.Body.Close()
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return rsp, nil
|
||||
}
|
||||
|
||||
func (t *Transport) roundTripOpt(req *http.Request, opt RoundTripOpt) (*http.Response, error) {
|
||||
t.initOnce.Do(func() { t.initErr = t.init() })
|
||||
if t.initErr != nil {
|
||||
return nil, t.initErr
|
||||
}
|
||||
|
||||
if req.URL == nil {
|
||||
return nil, errors.New("http3: nil Request.URL")
|
||||
}
|
||||
if req.URL.Scheme != "https" {
|
||||
return nil, fmt.Errorf("http3: unsupported protocol scheme: %s", req.URL.Scheme)
|
||||
}
|
||||
if req.URL.Host == "" {
|
||||
return nil, errors.New("http3: no Host in request URL")
|
||||
}
|
||||
if req.Header == nil {
|
||||
return nil, errors.New("http3: nil Request.Header")
|
||||
}
|
||||
if req.Method != "" && !validMethod(req.Method) {
|
||||
return nil, fmt.Errorf("http3: invalid method %q", req.Method)
|
||||
}
|
||||
for k, vv := range req.Header {
|
||||
// Skip validation for special header order keys
|
||||
if k == httpcommon.HeaderOrderKey || k == httpcommon.PHeaderOrderKey {
|
||||
continue
|
||||
}
|
||||
if !httpguts.ValidHeaderFieldName(k) {
|
||||
return nil, fmt.Errorf("http3: invalid http header field name %q", k)
|
||||
}
|
||||
for _, v := range vv {
|
||||
if !httpguts.ValidHeaderFieldValue(v) {
|
||||
return nil, fmt.Errorf("http3: invalid http header field value %q for key %v", v, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return t.doRoundTripOpt(req, opt, false)
|
||||
}
|
||||
|
||||
func (t *Transport) doRoundTripOpt(req *http.Request, opt RoundTripOpt, isRetried bool) (*http.Response, error) {
|
||||
hostname := authorityAddr(hostnameFromURL(req.URL))
|
||||
trace := httptrace.ContextClientTrace(req.Context())
|
||||
traceGetConn(trace, hostname)
|
||||
cl, isReused, err := t.getClient(req.Context(), hostname, opt.OnlyCachedConn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
select {
|
||||
case <-cl.dialing:
|
||||
case <-req.Context().Done():
|
||||
return nil, context.Cause(req.Context())
|
||||
}
|
||||
|
||||
if cl.dialErr != nil {
|
||||
t.removeClient(hostname)
|
||||
return nil, cl.dialErr
|
||||
}
|
||||
defer cl.useCount.Add(-1)
|
||||
traceGotConn(trace, cl.conn, isReused)
|
||||
rsp, err := cl.clientConn.RoundTrip(req)
|
||||
if err != nil {
|
||||
// request aborted due to context cancellation
|
||||
select {
|
||||
case <-req.Context().Done():
|
||||
return nil, err
|
||||
default:
|
||||
}
|
||||
if isRetried {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
t.removeClient(hostname)
|
||||
req, err = canRetryRequest(err, req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return t.doRoundTripOpt(req, opt, true)
|
||||
}
|
||||
return rsp, nil
|
||||
}
|
||||
|
||||
func canRetryRequest(err error, req *http.Request) (*http.Request, error) {
|
||||
// error occurred while opening the stream, we can be sure that the request wasn't sent out
|
||||
var connErr *errConnUnusable
|
||||
if errors.As(err, &connErr) {
|
||||
return req, nil
|
||||
}
|
||||
|
||||
// If the request stream is reset, we can only be sure that the request wasn't processed
|
||||
// if the error code is H3_REQUEST_REJECTED.
|
||||
var e *Error
|
||||
if !errors.As(err, &e) || e.ErrorCode != ErrCodeRequestRejected {
|
||||
return nil, err
|
||||
}
|
||||
// if the body is nil (or http.NoBody), it's safe to reuse this request and its body
|
||||
if req.Body == nil || req.Body == http.NoBody {
|
||||
return req, nil
|
||||
}
|
||||
// if the request body can be reset back to its original state via req.GetBody, do that
|
||||
if req.GetBody != nil {
|
||||
newBody, err := req.GetBody()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reqCopy := *req
|
||||
reqCopy.Body = newBody
|
||||
req = &reqCopy
|
||||
return &reqCopy, nil
|
||||
}
|
||||
return nil, fmt.Errorf("http3: Transport: cannot retry err [%w] after Request.Body was written; define Request.GetBody to avoid this error", err)
|
||||
}
|
||||
|
||||
// RoundTrip does a round trip.
|
||||
func (t *Transport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return t.RoundTripOpt(req, RoundTripOpt{})
|
||||
}
|
||||
|
||||
func (t *Transport) getClient(ctx context.Context, hostname string, onlyCached bool) (rtc *roundTripperWithCount, isReused bool, err error) {
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
if t.closed {
|
||||
return nil, false, ErrTransportClosed
|
||||
}
|
||||
|
||||
if t.clients == nil {
|
||||
t.clients = make(map[string]*roundTripperWithCount)
|
||||
}
|
||||
|
||||
cl, ok := t.clients[hostname]
|
||||
if !ok {
|
||||
if onlyCached {
|
||||
return nil, false, ErrNoCachedConn
|
||||
}
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
cl = &roundTripperWithCount{
|
||||
dialing: make(chan struct{}),
|
||||
cancel: cancel,
|
||||
}
|
||||
go func() {
|
||||
defer close(cl.dialing)
|
||||
defer cancel()
|
||||
conn, rt, err := t.dial(ctx, hostname)
|
||||
if err != nil {
|
||||
cl.dialErr = err
|
||||
return
|
||||
}
|
||||
cl.conn = conn
|
||||
cl.clientConn = rt
|
||||
}()
|
||||
t.clients[hostname] = cl
|
||||
}
|
||||
select {
|
||||
case <-cl.dialing:
|
||||
if cl.dialErr != nil {
|
||||
delete(t.clients, hostname)
|
||||
return nil, false, cl.dialErr
|
||||
}
|
||||
select {
|
||||
case <-cl.conn.HandshakeComplete():
|
||||
isReused = true
|
||||
default:
|
||||
}
|
||||
default:
|
||||
}
|
||||
cl.useCount.Add(1)
|
||||
return cl, isReused, nil
|
||||
}
|
||||
|
||||
func (t *Transport) dial(ctx context.Context, hostname string) (*quic.Conn, clientConn, error) {
|
||||
var tlsConf *tls.Config
|
||||
if t.TLSClientConfig == nil {
|
||||
tlsConf = &tls.Config{}
|
||||
} else {
|
||||
tlsConf = t.TLSClientConfig.Clone()
|
||||
}
|
||||
if tlsConf.ServerName == "" {
|
||||
sni, _, err := net.SplitHostPort(hostname)
|
||||
if err != nil {
|
||||
// It's ok if net.SplitHostPort returns an error - it could be a hostname/IP address without a port.
|
||||
sni = hostname
|
||||
}
|
||||
tlsConf.ServerName = sni
|
||||
}
|
||||
// Replace existing ALPNs by H3
|
||||
tlsConf.NextProtos = []string{NextProtoH3}
|
||||
|
||||
dial := t.Dial
|
||||
if dial == nil {
|
||||
dial = func(ctx context.Context, addr string, tlsCfg *tls.Config, cfg *quic.Config) (*quic.Conn, error) {
|
||||
network := "udp"
|
||||
udpAddr, err := t.resolveUDPAddr(ctx, network, addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
trace := httptrace.ContextClientTrace(ctx)
|
||||
traceConnectStart(trace, network, udpAddr.String())
|
||||
traceTLSHandshakeStart(trace)
|
||||
conn, err := t.transport.DialEarly(ctx, udpAddr, tlsCfg, cfg)
|
||||
var state tls.ConnectionState
|
||||
if conn != nil {
|
||||
state = conn.ConnectionState().TLS
|
||||
}
|
||||
traceTLSHandshakeDone(trace, state, err)
|
||||
traceConnectDone(trace, network, udpAddr.String(), err)
|
||||
return conn, err
|
||||
}
|
||||
}
|
||||
conn, err := dial(ctx, hostname, tlsConf, t.QUICConfig)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
clientConn := t.newClientConn(conn)
|
||||
go func() {
|
||||
for {
|
||||
str, err := conn.AcceptUniStream(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go clientConn.handleUnidirectionalStream(str)
|
||||
}
|
||||
}()
|
||||
return conn, clientConn, nil
|
||||
}
|
||||
|
||||
func (t *Transport) resolveUDPAddr(ctx context.Context, network, addr string) (*net.UDPAddr, error) {
|
||||
host, portStr, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
port, err := net.LookupPort(network, portStr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resolver := net.DefaultResolver
|
||||
ipAddrs, err := resolver.LookupIPAddr(ctx, host)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
addrs := addrList(ipAddrs)
|
||||
ip := addrs.forResolve(network, addr)
|
||||
return &net.UDPAddr{IP: ip.IP, Port: port, Zone: ip.Zone}, nil
|
||||
}
|
||||
|
||||
func (t *Transport) removeClient(hostname string) {
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
if t.clients == nil {
|
||||
return
|
||||
}
|
||||
delete(t.clients, hostname)
|
||||
}
|
||||
|
||||
// NewClientConn creates a new HTTP/3 client connection on top of a QUIC connection.
|
||||
// Most users should use RoundTrip instead of creating a connection directly.
|
||||
// Specifically, it is not needed to perform GET, POST, HEAD and CONNECT requests.
|
||||
//
|
||||
// Obtaining a ClientConn is only needed for more advanced use cases, such as
|
||||
// using Extended CONNECT for WebTransport or the various MASQUE protocols.
|
||||
func (t *Transport) NewClientConn(conn *quic.Conn) *ClientConn {
|
||||
c := newClientConn(
|
||||
conn,
|
||||
t.EnableDatagrams,
|
||||
t.AdditionalSettings,
|
||||
t.MaxResponseHeaderBytes,
|
||||
t.DisableCompression,
|
||||
t.Logger,
|
||||
)
|
||||
go func() {
|
||||
for {
|
||||
str, err := conn.AcceptUniStream(context.Background())
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
go c.handleUnidirectionalStream(str)
|
||||
}
|
||||
}()
|
||||
return c
|
||||
}
|
||||
|
||||
// NewRawClientConn creates a new low-level HTTP/3 client connection on top of a QUIC connection.
|
||||
// Unlike NewClientConn, the returned RawClientConn allows the application to take control
|
||||
// of the stream accept loops, by calling HandleUnidirectionalStream for incoming unidirectional
|
||||
// streams and HandleBidirectionalStream for incoming bidirectional streams.
|
||||
func (t *Transport) NewRawClientConn(conn *quic.Conn) *RawClientConn {
|
||||
return &RawClientConn{
|
||||
ClientConn: newClientConn(
|
||||
conn,
|
||||
t.EnableDatagrams,
|
||||
t.AdditionalSettings,
|
||||
t.MaxResponseHeaderBytes,
|
||||
t.DisableCompression,
|
||||
t.Logger,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
// Close closes the QUIC connections that this Transport has used.
|
||||
// A Transport cannot be used after it has been closed.
|
||||
func (t *Transport) Close() error {
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
for _, cl := range t.clients {
|
||||
if err := cl.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
t.clients = nil
|
||||
if t.transport != nil {
|
||||
if err := t.transport.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := t.transport.Conn.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
t.transport = nil
|
||||
}
|
||||
t.closed = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func hostnameFromURL(url *url.URL) string {
|
||||
if url != nil {
|
||||
return url.Host
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func validMethod(method string) bool {
|
||||
/*
|
||||
Method = "OPTIONS" ; Section 9.2
|
||||
| "GET" ; Section 9.3
|
||||
| "HEAD" ; Section 9.4
|
||||
| "POST" ; Section 9.5
|
||||
| "PUT" ; Section 9.6
|
||||
| "DELETE" ; Section 9.7
|
||||
| "TRACE" ; Section 9.8
|
||||
| "CONNECT" ; Section 9.9
|
||||
| extension-method
|
||||
extension-method = token
|
||||
token = 1*<any CHAR except CTLs or separators>
|
||||
*/
|
||||
return len(method) > 0 && strings.IndexFunc(method, isNotToken) == -1
|
||||
}
|
||||
|
||||
// copied from net/http/http.go
|
||||
func isNotToken(r rune) bool {
|
||||
return !httpguts.IsTokenRune(r)
|
||||
}
|
||||
|
||||
// CloseIdleConnections closes any QUIC connections in the transport's pool that are currently idle.
|
||||
// An idle connection is one that was previously used for requests but is now sitting unused.
|
||||
// This method does not interrupt any connections currently in use.
|
||||
// It also does not affect connections obtained via NewClientConn.
|
||||
func (t *Transport) CloseIdleConnections() {
|
||||
t.mutex.Lock()
|
||||
defer t.mutex.Unlock()
|
||||
for hostname, cl := range t.clients {
|
||||
if cl.useCount.Load() == 0 {
|
||||
cl.Close()
|
||||
delete(t.clients, hostname)
|
||||
}
|
||||
}
|
||||
}
|
||||
+58
@@ -0,0 +1,58 @@
|
||||
// Package retryafter parses values of the HTTP Retry-After header
|
||||
// as defined by RFC 7231 §7.1.3.
|
||||
//
|
||||
// Two forms are recognised: an integer count of seconds (delay-seconds)
|
||||
// and an HTTP-date in any of the three formats accepted by RFC 7231
|
||||
// (IMF-fixdate, RFC 850, ANSI C asctime). The HTTP-date branch is
|
||||
// handled by github.com/enetx/http.ParseTime, which mirrors the
|
||||
// standard library's net/http.ParseTime.
|
||||
package retryafter
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/enetx/http"
|
||||
)
|
||||
|
||||
// Parse interprets value as a Retry-After header field.
|
||||
//
|
||||
// On success it returns the wait duration and true. A successful parse of "0",
|
||||
// a past HTTP-date, or any negative time.Until(parsed) is clamped to zero so
|
||||
// that callers using max(retryWait, parsed) preserve their minimum floor.
|
||||
//
|
||||
// For missing, empty, malformed, fractional, or negative integer values it
|
||||
// returns (0, false) so callers fall back to their own wait policy.
|
||||
//
|
||||
// The now argument is injected to keep HTTP-date arithmetic deterministic in
|
||||
// tests; production callers pass time.Now() immediately before scheduling the
|
||||
// timer.
|
||||
//
|
||||
// Multiple Retry-After headers are undefined by the RFC; callers must select
|
||||
// a single value via http.Header.Get before calling Parse.
|
||||
func Parse(value string, now time.Time) (time.Duration, bool) {
|
||||
v := strings.TrimSpace(value)
|
||||
if v == "" {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
if n, err := strconv.Atoi(v); err == nil {
|
||||
if n < 0 {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
return time.Duration(n) * time.Second, true
|
||||
}
|
||||
|
||||
if t, err := http.ParseTime(v); err == nil {
|
||||
d := t.Sub(now)
|
||||
if d < 0 {
|
||||
return 0, true
|
||||
}
|
||||
|
||||
return d, true
|
||||
}
|
||||
|
||||
return 0, false
|
||||
}
|
||||
+209
@@ -0,0 +1,209 @@
|
||||
package surf
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"mime/multipart"
|
||||
"net/textproto"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/g/fs"
|
||||
)
|
||||
|
||||
// Multipart represents multipart form data with fields and files.
|
||||
type Multipart struct {
|
||||
fields g.MapOrd[g.String, g.String]
|
||||
files g.Slice[*MultipartFile]
|
||||
retry bool
|
||||
}
|
||||
|
||||
// MultipartFile represents a single file for multipart upload.
|
||||
type MultipartFile struct {
|
||||
fieldName g.String
|
||||
fileName g.String
|
||||
contentType g.String
|
||||
file *fs.File
|
||||
reader io.Reader
|
||||
}
|
||||
|
||||
// NewMultipart creates a new empty Multipart object.
|
||||
func NewMultipart() *Multipart {
|
||||
return &Multipart{
|
||||
fields: g.NewMapOrd[g.String, g.String](),
|
||||
files: g.NewSlice[*MultipartFile](),
|
||||
}
|
||||
}
|
||||
|
||||
// Field adds a form field to the multipart.
|
||||
func (m *Multipart) Field(name, value g.String) *Multipart {
|
||||
m.fields.Insert(name, value)
|
||||
return m
|
||||
}
|
||||
|
||||
// File adds a physical file to the multipart.
|
||||
func (m *Multipart) File(fieldName g.String, file *fs.File) *Multipart {
|
||||
f := &MultipartFile{
|
||||
fieldName: fieldName,
|
||||
fileName: file.Name(),
|
||||
file: file,
|
||||
}
|
||||
|
||||
m.files.Push(f)
|
||||
return m
|
||||
}
|
||||
|
||||
// FileReader adds a file from io.Reader to the multipart.
|
||||
func (m *Multipart) FileReader(fieldName, fileName g.String, reader io.Reader) *Multipart {
|
||||
f := &MultipartFile{
|
||||
fieldName: fieldName,
|
||||
fileName: fileName,
|
||||
reader: reader,
|
||||
}
|
||||
|
||||
m.files.Push(f)
|
||||
return m
|
||||
}
|
||||
|
||||
// FileString adds a file from string content to the multipart.
|
||||
func (m *Multipart) FileString(fieldName, fileName, content g.String) *Multipart {
|
||||
f := &MultipartFile{
|
||||
fieldName: fieldName,
|
||||
fileName: fileName,
|
||||
reader: content.Reader(),
|
||||
}
|
||||
|
||||
m.files.Push(f)
|
||||
return m
|
||||
}
|
||||
|
||||
// FileBytes adds a file from byte slice to the multipart.
|
||||
func (m *Multipart) FileBytes(fieldName, fileName g.String, data g.Bytes) *Multipart {
|
||||
f := &MultipartFile{
|
||||
fieldName: fieldName,
|
||||
fileName: fileName,
|
||||
reader: data.Reader(),
|
||||
}
|
||||
|
||||
m.files.Push(f)
|
||||
return m
|
||||
}
|
||||
|
||||
// ContentType sets the content type for the last added file.
|
||||
// Must be called immediately after File/FileReader/FileString/FileBytes.
|
||||
func (m *Multipart) ContentType(ct g.String) *Multipart {
|
||||
if last := m.files.Last(); last.IsSome() {
|
||||
last.Some().contentType = ct
|
||||
}
|
||||
|
||||
return m
|
||||
}
|
||||
|
||||
// FileName overrides the filename for the last added file.
|
||||
// Useful when you want a different name than the physical file.
|
||||
func (m *Multipart) FileName(name g.String) *Multipart {
|
||||
if last := m.files.Last(); last.IsSome() {
|
||||
last.Some().fileName = name
|
||||
}
|
||||
|
||||
return m
|
||||
}
|
||||
|
||||
// Retry controls whether the multipart body should be buffered in memory
|
||||
// to support retries on status codes (429, 503, 5xx, etc.).
|
||||
//
|
||||
// When set, the body is fully read into memory before sending,
|
||||
// allowing the client to replay it on retry.
|
||||
//
|
||||
// Recommended only for small requests (≤ 5–10 MB).
|
||||
func (m *Multipart) Retry() *Multipart {
|
||||
m.retry = true
|
||||
return m
|
||||
}
|
||||
|
||||
// prepareWriter writes the multipart data to a writer and returns the content type and write error.
|
||||
func (m *Multipart) prepareWriter(boundary func() g.String) (io.ReadCloser, string, error) {
|
||||
pr, pw := io.Pipe()
|
||||
writer := multipart.NewWriter(pw)
|
||||
|
||||
if boundary != nil {
|
||||
if err := writer.SetBoundary(boundary().Std()); err != nil {
|
||||
_ = pw.CloseWithError(err)
|
||||
_ = pr.Close()
|
||||
return nil, "", err
|
||||
}
|
||||
}
|
||||
|
||||
go func() {
|
||||
err := func() error {
|
||||
for key, val := range m.fields.Iter() {
|
||||
part, err := writer.CreateFormField(key.Std())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := io.Copy(part, val.Reader()); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
for file := range m.files.Iter() {
|
||||
var reader io.Reader
|
||||
|
||||
if file.file != nil {
|
||||
res := file.file.Open()
|
||||
if res.IsErr() {
|
||||
return fmt.Errorf("cannot open file %q: %w", file.file.Name(), res.Err())
|
||||
}
|
||||
|
||||
opened := res.Ok()
|
||||
defer opened.Close()
|
||||
|
||||
reader = opened.Std()
|
||||
} else if file.reader != nil {
|
||||
reader = file.reader
|
||||
} else {
|
||||
return fmt.Errorf("multipart file %q has no content source", file.fileName.Std())
|
||||
}
|
||||
|
||||
ct := file.contentType.Std()
|
||||
if ct == "" {
|
||||
ext := filepath.Ext(file.fileName.Std())
|
||||
ct = mime.TypeByExtension(ext)
|
||||
if ct == "" {
|
||||
ct = "application/octet-stream"
|
||||
}
|
||||
}
|
||||
|
||||
disposition := fmt.Sprintf(
|
||||
`form-data; name="%s"; filename="%s"`,
|
||||
escapeQuotes(file.fieldName),
|
||||
escapeQuotes(file.fileName),
|
||||
)
|
||||
|
||||
h := textproto.MIMEHeader{
|
||||
"Content-Disposition": {disposition},
|
||||
"Content-Type": {ct},
|
||||
}
|
||||
|
||||
part, err := writer.CreatePart(h)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := io.Copy(part, reader); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return writer.Close()
|
||||
}()
|
||||
|
||||
pw.CloseWithError(err)
|
||||
}()
|
||||
|
||||
return pr, writer.FormDataContentType(), nil
|
||||
}
|
||||
|
||||
func escapeQuotes(s g.String) string { return s.ReplaceMulti(`\`, `\\`, `"`, `\"`).Std() }
|
||||
+48
@@ -0,0 +1,48 @@
|
||||
package socks4
|
||||
|
||||
import "fmt"
|
||||
|
||||
type (
|
||||
ErrWrongNetwork struct{}
|
||||
ErrConnRejected struct{}
|
||||
ErrIdentRequired struct{}
|
||||
ErrDialFailed struct{ err error }
|
||||
ErrBuffer struct{ err error }
|
||||
ErrIO struct{ err error }
|
||||
ErrInvalidResponse struct{ resp byte }
|
||||
|
||||
ErrWrongAddr struct {
|
||||
msg string
|
||||
err error
|
||||
}
|
||||
|
||||
ErrHostUnknown struct {
|
||||
msg string
|
||||
err error
|
||||
}
|
||||
)
|
||||
|
||||
func (e *ErrDialFailed) Error() string { return fmt.Sprintf("socks4 dial %v", e.err) }
|
||||
func (e *ErrDialFailed) Unwrap() error { return e.err }
|
||||
|
||||
func (e *ErrHostUnknown) Error() string {
|
||||
return fmt.Sprintf("unable to find IP address of host %s", e.msg)
|
||||
}
|
||||
func (e *ErrHostUnknown) Unwrap() error { return e.err }
|
||||
|
||||
func (e *ErrBuffer) Error() string { return "unable write into buffer" }
|
||||
func (e *ErrBuffer) Unwrap() error { return e.err }
|
||||
|
||||
func (e *ErrIO) Error() string { return "io error" }
|
||||
func (e *ErrIO) Unwrap() error { return e.err }
|
||||
|
||||
func (e *ErrWrongAddr) Error() string { return fmt.Sprintf("wrong addr: %s, error: %v", e.msg, e.err) }
|
||||
func (e *ErrWrongAddr) Unwrap() error { return e.err }
|
||||
|
||||
func (e *ErrWrongNetwork) Error() string { return "network should be tcp or tcp4" }
|
||||
|
||||
func (e *ErrConnRejected) Error() string { return "connection to remote host was rejected" }
|
||||
func (e *ErrIdentRequired) Error() string { return "valid ident required" }
|
||||
func (e *ErrInvalidResponse) Error() string {
|
||||
return fmt.Sprintf("unknown socks4 server response 0x%02x", e.resp)
|
||||
}
|
||||
+207
@@ -0,0 +1,207 @@
|
||||
package socks4
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"net"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"golang.org/x/net/proxy"
|
||||
)
|
||||
|
||||
const (
|
||||
socksVersion = 0x04
|
||||
socksConnect = 0x01
|
||||
socksBind = 0x02
|
||||
|
||||
accessGranted = 0x5a
|
||||
accessRejected = 0x5b
|
||||
accessIdentRequired = 0x5c
|
||||
accessIdentFailed = 0x5d
|
||||
|
||||
minRequestLen = 8
|
||||
)
|
||||
|
||||
var Ident = "nobody@0.0.0.0"
|
||||
|
||||
func init() {
|
||||
proxy.RegisterDialerType("socks4", func(u *url.URL, d proxy.Dialer) (proxy.Dialer, error) {
|
||||
return socks4{url: u, dialer: d}, nil
|
||||
})
|
||||
|
||||
proxy.RegisterDialerType("socks4a", func(u *url.URL, d proxy.Dialer) (proxy.Dialer, error) {
|
||||
return socks4{url: u, dialer: d}, nil
|
||||
})
|
||||
}
|
||||
|
||||
type socks4 struct {
|
||||
url *url.URL
|
||||
dialer proxy.Dialer
|
||||
}
|
||||
|
||||
// DialContext implements proxy.ContextDialer interface
|
||||
func (s socks4) DialContext(ctx context.Context, network, addr string) (c net.Conn, err error) {
|
||||
if network != "tcp" && network != "tcp4" {
|
||||
return nil, new(ErrWrongNetwork)
|
||||
}
|
||||
|
||||
// Use context-aware dialer if available
|
||||
if cd, ok := s.dialer.(proxy.ContextDialer); ok {
|
||||
c, err = cd.DialContext(ctx, network, s.url.Host)
|
||||
} else {
|
||||
c, err = s.dialer.Dial(network, s.url.Host)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, &ErrDialFailed{err}
|
||||
}
|
||||
|
||||
// close connection later if we got an error
|
||||
defer func() {
|
||||
if err != nil && c != nil {
|
||||
_ = c.Close()
|
||||
}
|
||||
}()
|
||||
|
||||
// Set deadline from context
|
||||
if deadline, ok := ctx.Deadline(); ok {
|
||||
c.SetDeadline(deadline)
|
||||
defer c.SetDeadline(time.Time{})
|
||||
}
|
||||
|
||||
// Check context before handshake
|
||||
if ctx.Err() != nil {
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
host, port, err := s.parseAddr(addr)
|
||||
if err != nil {
|
||||
return nil, &ErrWrongAddr{addr, err}
|
||||
}
|
||||
|
||||
ip := net.IPv4(0, 0, 0, 1)
|
||||
if !s.isSocks4a() {
|
||||
if ip, err = s.lookupAddr(ctx, host); err != nil {
|
||||
return nil, &ErrHostUnknown{host, err}
|
||||
}
|
||||
}
|
||||
|
||||
req, err := request{Host: host, Port: port, IP: ip, Is4a: s.isSocks4a()}.Bytes()
|
||||
if err != nil {
|
||||
return nil, &ErrBuffer{err}
|
||||
}
|
||||
|
||||
var i int
|
||||
i, err = c.Write(req)
|
||||
if err != nil {
|
||||
return c, &ErrIO{err}
|
||||
} else if i < minRequestLen {
|
||||
return c, &ErrIO{io.ErrShortWrite}
|
||||
}
|
||||
|
||||
var resp [8]byte
|
||||
i, err = c.Read(resp[:])
|
||||
if err != nil && err != io.EOF {
|
||||
return c, &ErrIO{err}
|
||||
} else if i != 8 {
|
||||
return c, &ErrIO{io.ErrUnexpectedEOF}
|
||||
}
|
||||
|
||||
switch resp[1] {
|
||||
case accessGranted:
|
||||
return c, nil
|
||||
case accessIdentRequired, accessIdentFailed:
|
||||
return c, new(ErrIdentRequired)
|
||||
case accessRejected:
|
||||
return c, new(ErrConnRejected)
|
||||
default:
|
||||
return c, &ErrInvalidResponse{resp[1]}
|
||||
}
|
||||
}
|
||||
|
||||
// Dial implements proxy.Dialer interface
|
||||
func (s socks4) Dial(network, addr string) (net.Conn, error) {
|
||||
return s.DialContext(context.Background(), network, addr)
|
||||
}
|
||||
|
||||
func (s socks4) lookupAddr(ctx context.Context, host string) (net.IP, error) {
|
||||
resolver := net.DefaultResolver
|
||||
ips, err := resolver.LookupIPAddr(ctx, host)
|
||||
if err != nil {
|
||||
return net.IP{}, err
|
||||
}
|
||||
|
||||
for _, ip := range ips {
|
||||
if v4 := ip.IP.To4(); v4 != nil {
|
||||
return v4, nil
|
||||
}
|
||||
}
|
||||
|
||||
return net.IP{}, &net.DNSError{Err: "no IPv4 address", Name: host}
|
||||
}
|
||||
|
||||
func (s socks4) isSocks4a() bool {
|
||||
return s.url.Scheme == "socks4a"
|
||||
}
|
||||
|
||||
func (s socks4) parseAddr(addr string) (host string, iport int, err error) {
|
||||
var port string
|
||||
|
||||
host, port, err = net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
|
||||
iport, err = strconv.Atoi(port)
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
|
||||
return host, iport, err
|
||||
}
|
||||
|
||||
type request struct {
|
||||
Host string
|
||||
Port int
|
||||
IP net.IP
|
||||
Is4a bool
|
||||
|
||||
err error
|
||||
buf bytes.Buffer
|
||||
}
|
||||
|
||||
func (r *request) write(b []byte) {
|
||||
if r.err == nil {
|
||||
_, r.err = r.buf.Write(b)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *request) writeString(s string) {
|
||||
if r.err == nil {
|
||||
_, r.err = r.buf.WriteString(s)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *request) writeBigEndian(data any) {
|
||||
if r.err == nil {
|
||||
r.err = binary.Write(&r.buf, binary.BigEndian, data)
|
||||
}
|
||||
}
|
||||
|
||||
func (r request) Bytes() ([]byte, error) {
|
||||
r.write([]byte{socksVersion, socksConnect})
|
||||
r.writeBigEndian(uint16(r.Port))
|
||||
r.writeBigEndian(r.IP.To4())
|
||||
r.writeString(Ident)
|
||||
r.write([]byte{0})
|
||||
if r.Is4a {
|
||||
r.writeString(r.Host)
|
||||
r.write([]byte{0})
|
||||
}
|
||||
|
||||
return r.buf.Bytes(), r.err
|
||||
}
|
||||
+134
@@ -0,0 +1,134 @@
|
||||
package surf
|
||||
|
||||
import (
|
||||
"compress/gzip"
|
||||
"errors"
|
||||
"io"
|
||||
"sync"
|
||||
|
||||
"github.com/andybalholm/brotli"
|
||||
"github.com/enetx/g"
|
||||
"github.com/klauspost/compress/zstd"
|
||||
)
|
||||
|
||||
// decodedReadCloser owns both the decoder and the encoded source it wraps.
|
||||
// Closing only the decoder leaks the HTTP response body (and its deadline
|
||||
// goroutine); nesting this type also preserves ownership for encoding chains.
|
||||
type decodedReadCloser struct {
|
||||
decoder io.ReadCloser
|
||||
source io.ReadCloser
|
||||
}
|
||||
|
||||
func (reader *decodedReadCloser) Read(buffer []byte) (int, error) {
|
||||
return reader.decoder.Read(buffer)
|
||||
}
|
||||
|
||||
func (reader *decodedReadCloser) Close() error {
|
||||
return errors.Join(reader.decoder.Close(), reader.source.Close())
|
||||
}
|
||||
|
||||
var (
|
||||
// zstdDecoderPool pools zstd.Decoder instances.
|
||||
zstdDecoderPool = sync.Pool{
|
||||
New: func() any {
|
||||
dec, _ := zstd.NewReader(nil)
|
||||
return dec
|
||||
},
|
||||
}
|
||||
|
||||
// gzipReaderPool pools gzip.Reader instances.
|
||||
gzipReaderPool = sync.Pool{
|
||||
New: func() any {
|
||||
return new(gzip.Reader)
|
||||
},
|
||||
}
|
||||
|
||||
// brotliReaderPool pools brotli.Reader instances.
|
||||
brotliReaderPool = sync.Pool{
|
||||
New: func() any {
|
||||
return brotli.NewReader(nil)
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
// zstdReadCloser wraps a zstd decoder and returns it to the pool on Close.
|
||||
type zstdReadCloser struct {
|
||||
dec *zstd.Decoder
|
||||
}
|
||||
|
||||
// Read reads decompressed data from the decoder.
|
||||
func (zr *zstdReadCloser) Read(p []byte) (int, error) {
|
||||
return zr.dec.Read(p)
|
||||
}
|
||||
|
||||
// Close resets the decoder and returns it to the pool.
|
||||
func (zr *zstdReadCloser) Close() error {
|
||||
if err := zr.dec.Reset(nil); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
zstdDecoderPool.Put(zr.dec)
|
||||
return nil
|
||||
}
|
||||
|
||||
// gzipReadCloser wraps a gzip reader and returns it to the pool on Close.
|
||||
type gzipReadCloser struct {
|
||||
*gzip.Reader
|
||||
}
|
||||
|
||||
// Close closes the reader and returns it to the pool.
|
||||
func (gr *gzipReadCloser) Close() error {
|
||||
if err := gr.Reader.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
gzipReaderPool.Put(gr.Reader)
|
||||
return nil
|
||||
}
|
||||
|
||||
// brotliReadCloser wraps a brotli reader and returns it to the pool on Close.
|
||||
type brotliReadCloser struct {
|
||||
*brotli.Reader
|
||||
}
|
||||
|
||||
// Close returns the reader to the pool.
|
||||
func (br *brotliReadCloser) Close() error {
|
||||
brotliReaderPool.Put(br.Reader)
|
||||
return nil
|
||||
}
|
||||
|
||||
// acquireGzipReader gets a gzip.Reader from the pool and resets it with the provided reader.
|
||||
// Returns the reader wrapped in gzipReadCloser for automatic pool management.
|
||||
func acquireGzipReader(r io.Reader) g.Result[io.ReadCloser] {
|
||||
gr := gzipReaderPool.Get().(*gzip.Reader)
|
||||
if err := gr.Reset(r); err != nil {
|
||||
gzipReaderPool.Put(gr)
|
||||
return g.Err[io.ReadCloser](err)
|
||||
}
|
||||
|
||||
return g.Ok[io.ReadCloser](&gzipReadCloser{Reader: gr})
|
||||
}
|
||||
|
||||
// acquireBrotliReader gets a brotli.Reader from the pool and resets it with the provided reader.
|
||||
// Returns the reader wrapped in brotliReadCloser for automatic pool management.
|
||||
func acquireBrotliReader(r io.Reader) g.Result[io.ReadCloser] {
|
||||
br := brotliReaderPool.Get().(*brotli.Reader)
|
||||
if err := br.Reset(r); err != nil {
|
||||
brotliReaderPool.Put(br)
|
||||
return g.Err[io.ReadCloser](err)
|
||||
}
|
||||
|
||||
return g.Ok[io.ReadCloser](&brotliReadCloser{Reader: br})
|
||||
}
|
||||
|
||||
// acquireZstdReader gets a zstd.Decoder from the pool and resets it with the provided reader.
|
||||
// Returns the decoder wrapped in zstdReadCloser for automatic pool management.
|
||||
func acquireZstdReader(r io.Reader) g.Result[io.ReadCloser] {
|
||||
dec := zstdDecoderPool.Get().(*zstd.Decoder)
|
||||
if err := dec.Reset(r); err != nil {
|
||||
zstdDecoderPool.Put(dec)
|
||||
return g.Err[io.ReadCloser](err)
|
||||
}
|
||||
|
||||
return g.Ok[io.ReadCloser](&zstdReadCloser{dec: dec})
|
||||
}
|
||||
+109
@@ -0,0 +1,109 @@
|
||||
package surf
|
||||
|
||||
import (
|
||||
"github.com/enetx/http2"
|
||||
"github.com/enetx/surf/profiles"
|
||||
)
|
||||
|
||||
// h2adapter wraps *HTTP2Settings to satisfy profiles.H2Config (which returns the interface type
|
||||
// instead of *HTTP2Settings, so direct method satisfaction is impossible). Each method delegates
|
||||
// to the underlying *HTTP2Settings and returns the adapter to keep the chain fluent.
|
||||
type h2adapter struct{ s *HTTP2Settings }
|
||||
|
||||
func (a h2adapter) HeaderTableSize(v uint32) profiles.H2Config {
|
||||
a.s.HeaderTableSize(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) EnablePush(v uint32) profiles.H2Config {
|
||||
a.s.EnablePush(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) MaxConcurrentStreams(v uint32) profiles.H2Config {
|
||||
a.s.MaxConcurrentStreams(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) InitialWindowSize(v uint32) profiles.H2Config {
|
||||
a.s.InitialWindowSize(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) MaxFrameSize(v uint32) profiles.H2Config {
|
||||
a.s.MaxFrameSize(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) MaxHeaderListSize(v uint32) profiles.H2Config {
|
||||
a.s.MaxHeaderListSize(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) NoRFC7540Priorities(v uint32) profiles.H2Config {
|
||||
a.s.NoRFC7540Priorities(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) ConnectionFlow(v uint32) profiles.H2Config {
|
||||
a.s.ConnectionFlow(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) InitialStreamID(v uint32) profiles.H2Config {
|
||||
a.s.InitialStreamID(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) PriorityParam(v http2.PriorityParam) profiles.H2Config {
|
||||
a.s.PriorityParam(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h2adapter) PriorityFrames(v []http2.PriorityFrame) profiles.H2Config {
|
||||
a.s.PriorityFrames(v)
|
||||
return a
|
||||
}
|
||||
|
||||
// h3adapter wraps *HTTP3Settings to satisfy profiles.H3Config.
|
||||
type h3adapter struct{ s *HTTP3Settings }
|
||||
|
||||
func (a h3adapter) QpackMaxTableCapacity(v uint64) profiles.H3Config {
|
||||
a.s.QpackMaxTableCapacity(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h3adapter) MaxFieldSectionSize(v uint64) profiles.H3Config {
|
||||
a.s.MaxFieldSectionSize(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h3adapter) QpackBlockedStreams(v uint64) profiles.H3Config {
|
||||
a.s.QpackBlockedStreams(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h3adapter) EnableConnectProtocol(v uint64) profiles.H3Config {
|
||||
a.s.EnableConnectProtocol(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h3adapter) SettingsH3Datagram(v uint64) profiles.H3Config {
|
||||
a.s.SettingsH3Datagram(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h3adapter) H3Datagram(v uint64) profiles.H3Config {
|
||||
a.s.H3Datagram(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h3adapter) EnableWebtransport(v uint64) profiles.H3Config {
|
||||
a.s.EnableWebtransport(v)
|
||||
return a
|
||||
}
|
||||
|
||||
func (a h3adapter) Grease() profiles.H3Config {
|
||||
a.s.Grease()
|
||||
return a
|
||||
}
|
||||
+81
@@ -0,0 +1,81 @@
|
||||
package chrome
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
|
||||
"github.com/enetx/g"
|
||||
)
|
||||
|
||||
// Blink implementation: https://source.chromium.org/chromium/chromium/src/+/main:third_party/blink/renderer/platform/network/form_data_encoder.cc;drc=1d694679493c7b2f7b9df00e967b4f8699321093;l=130
|
||||
// WebKit implementation: https://github.com/WebKit/WebKit/blob/main/Source/WebCore/platform/network/FormDataBuilder.cpp#L120
|
||||
func Boundary() g.String {
|
||||
// C++
|
||||
// Vector<uint8_t> generateUniqueBoundaryString()
|
||||
// {
|
||||
// Vector<uint8_t> boundary;
|
||||
//
|
||||
// // The RFC 2046 spec says the alphanumeric characters plus the
|
||||
// // following characters are legal for boundaries: '()+_,-./:=?
|
||||
// // However the following characters, though legal, cause some sites
|
||||
// // to fail: (),./:=+
|
||||
// // Note that our algorithm makes it twice as much likely for 'A' or 'B'
|
||||
// // to appear in the boundary string, because 0x41 and 0x42 are present in
|
||||
// // the below array twice.
|
||||
// static constexpr std::array<char, 64> alphaNumericEncodingMap {
|
||||
// 0x41, 0x42, 0x43, 0x44, 0x45, 0x46, 0x47, 0x48,
|
||||
// 0x49, 0x4A, 0x4B, 0x4C, 0x4D, 0x4E, 0x4F, 0x50,
|
||||
// 0x51, 0x52, 0x53, 0x54, 0x55, 0x56, 0x57, 0x58,
|
||||
// 0x59, 0x5A, 0x61, 0x62, 0x63, 0x64, 0x65, 0x66,
|
||||
// 0x67, 0x68, 0x69, 0x6A, 0x6B, 0x6C, 0x6D, 0x6E,
|
||||
// 0x6F, 0x70, 0x71, 0x72, 0x73, 0x74, 0x75, 0x76,
|
||||
// 0x77, 0x78, 0x79, 0x7A, 0x30, 0x31, 0x32, 0x33,
|
||||
// 0x34, 0x35, 0x36, 0x37, 0x38, 0x39, 0x41, 0x42
|
||||
// };
|
||||
//
|
||||
// // Start with an informative prefix.
|
||||
// append(boundary, "----WebKitFormBoundary");
|
||||
//
|
||||
// // Append 16 random 7-bit ASCII alphanumeric characters.
|
||||
// for (unsigned i = 0; i < 4; ++i) {
|
||||
// unsigned randomness = cryptographicallyRandomNumber<unsigned>();
|
||||
// boundary.append(alphaNumericEncodingMap[(randomness >> 24) & 0x3F]);
|
||||
// boundary.append(alphaNumericEncodingMap[(randomness >> 16) & 0x3F]);
|
||||
// boundary.append(alphaNumericEncodingMap[(randomness >> 8) & 0x3F]);
|
||||
// boundary.append(alphaNumericEncodingMap[randomness & 0x3F]);
|
||||
// }
|
||||
//
|
||||
// return boundary;
|
||||
// }
|
||||
|
||||
prefix := "----WebKitFormBoundary"
|
||||
|
||||
alphaNumericEncodingMap := []byte{
|
||||
0x41, 0x42, 0x43, 0x44, 0x45, 0x46, 0x47, 0x48,
|
||||
0x49, 0x4A, 0x4B, 0x4C, 0x4D, 0x4E, 0x4F, 0x50,
|
||||
0x51, 0x52, 0x53, 0x54, 0x55, 0x56, 0x57, 0x58,
|
||||
0x59, 0x5A, 0x61, 0x62, 0x63, 0x64, 0x65, 0x66,
|
||||
0x67, 0x68, 0x69, 0x6A, 0x6B, 0x6C, 0x6D, 0x6E,
|
||||
0x6F, 0x70, 0x71, 0x72, 0x73, 0x74, 0x75, 0x76,
|
||||
0x77, 0x78, 0x79, 0x7A, 0x30, 0x31, 0x32, 0x33,
|
||||
0x34, 0x35, 0x36, 0x37, 0x38, 0x39, 0x41, 0x42,
|
||||
}
|
||||
|
||||
boundary := []byte(prefix)
|
||||
|
||||
for range 4 {
|
||||
randomBytes := make([]byte, 4)
|
||||
rand.Read(randomBytes)
|
||||
|
||||
randomness := uint32(randomBytes[0])<<24 |
|
||||
uint32(randomBytes[1])<<16 |
|
||||
uint32(randomBytes[2])<<8 |
|
||||
uint32(randomBytes[3])
|
||||
|
||||
boundary = append(boundary, alphaNumericEncodingMap[(randomness>>24)&0x3F])
|
||||
boundary = append(boundary, alphaNumericEncodingMap[(randomness>>16)&0x3F])
|
||||
boundary = append(boundary, alphaNumericEncodingMap[(randomness>>8)&0x3F])
|
||||
boundary = append(boundary, alphaNumericEncodingMap[randomness&0x3F])
|
||||
}
|
||||
|
||||
return g.String(boundary)
|
||||
}
|
||||
+324
@@ -0,0 +1,324 @@
|
||||
package chrome
|
||||
|
||||
import (
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/http"
|
||||
"github.com/enetx/surf/header"
|
||||
"github.com/enetx/surf/profiles"
|
||||
)
|
||||
|
||||
// --- Header order maps -------------------------------------------------------
|
||||
|
||||
var headerOrderDesktop = g.Map[string, g.Slice[string]]{
|
||||
http.MethodGet: {
|
||||
":method",
|
||||
":authority",
|
||||
":scheme",
|
||||
":path",
|
||||
header.SEC_CH_UA,
|
||||
header.SEC_CH_UA_MOBILE,
|
||||
header.SEC_CH_UA_PLATFORM,
|
||||
header.AUTHORIZATION,
|
||||
header.UPGRADE_INSECURE_REQUESTS,
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_USER,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.REFERER,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.COOKIE,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodGet + "http3": {
|
||||
":method",
|
||||
":authority",
|
||||
":scheme",
|
||||
":path",
|
||||
header.SEC_CH_UA,
|
||||
header.SEC_CH_UA_MOBILE,
|
||||
header.SEC_CH_UA_PLATFORM,
|
||||
header.AUTHORIZATION,
|
||||
header.UPGRADE_INSECURE_REQUESTS,
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_USER,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.REFERER,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.COOKIE,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodPost: {
|
||||
":method",
|
||||
":authority",
|
||||
":scheme",
|
||||
":path",
|
||||
header.CONTENT_LENGTH,
|
||||
header.PRAGMA,
|
||||
header.CACHE_CONTROL,
|
||||
header.SEC_CH_UA_PLATFORM,
|
||||
header.AUTHORIZATION,
|
||||
header.USER_AGENT,
|
||||
header.SEC_CH_UA,
|
||||
header.CONTENT_TYPE,
|
||||
header.SEC_CH_UA_MOBILE,
|
||||
header.ACCEPT,
|
||||
header.ORIGIN,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.REFERER,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.COOKIE,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodPost + "http3": {
|
||||
":method",
|
||||
":authority",
|
||||
":scheme",
|
||||
":path",
|
||||
header.CONTENT_LENGTH,
|
||||
header.PRAGMA,
|
||||
header.CACHE_CONTROL,
|
||||
header.SEC_CH_UA_PLATFORM,
|
||||
header.AUTHORIZATION,
|
||||
header.USER_AGENT,
|
||||
header.SEC_CH_UA,
|
||||
header.CONTENT_TYPE,
|
||||
header.SEC_CH_UA_MOBILE,
|
||||
header.ACCEPT,
|
||||
header.ORIGIN,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.REFERER,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.COOKIE,
|
||||
header.PRIORITY,
|
||||
},
|
||||
}
|
||||
|
||||
// headerOrderMobile is a placeholder mobile variant. On the day real Chrome Android header
|
||||
// ordering is observed, this map is the single point to substitute it without touching desktop.
|
||||
// The literal is a physical copy of headerOrderDesktop so the two maps can diverge independently.
|
||||
var headerOrderMobile = g.Map[string, g.Slice[string]]{
|
||||
http.MethodGet: {
|
||||
":method",
|
||||
":authority",
|
||||
":scheme",
|
||||
":path",
|
||||
header.SEC_CH_UA,
|
||||
header.SEC_CH_UA_MOBILE,
|
||||
header.SEC_CH_UA_PLATFORM,
|
||||
header.AUTHORIZATION,
|
||||
header.UPGRADE_INSECURE_REQUESTS,
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_USER,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.REFERER,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.COOKIE,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodGet + "http3": {
|
||||
":method",
|
||||
":authority",
|
||||
":scheme",
|
||||
":path",
|
||||
header.SEC_CH_UA,
|
||||
header.SEC_CH_UA_MOBILE,
|
||||
header.SEC_CH_UA_PLATFORM,
|
||||
header.AUTHORIZATION,
|
||||
header.UPGRADE_INSECURE_REQUESTS,
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_USER,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.REFERER,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.COOKIE,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodPost: {
|
||||
":method",
|
||||
":authority",
|
||||
":scheme",
|
||||
":path",
|
||||
header.CONTENT_LENGTH,
|
||||
header.PRAGMA,
|
||||
header.CACHE_CONTROL,
|
||||
header.SEC_CH_UA_PLATFORM,
|
||||
header.AUTHORIZATION,
|
||||
header.USER_AGENT,
|
||||
header.SEC_CH_UA,
|
||||
header.CONTENT_TYPE,
|
||||
header.SEC_CH_UA_MOBILE,
|
||||
header.ACCEPT,
|
||||
header.ORIGIN,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.REFERER,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.COOKIE,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodPost + "http3": {
|
||||
":method",
|
||||
":authority",
|
||||
":scheme",
|
||||
":path",
|
||||
header.CONTENT_LENGTH,
|
||||
header.PRAGMA,
|
||||
header.CACHE_CONTROL,
|
||||
header.SEC_CH_UA_PLATFORM,
|
||||
header.AUTHORIZATION,
|
||||
header.USER_AGENT,
|
||||
header.SEC_CH_UA,
|
||||
header.CONTENT_TYPE,
|
||||
header.SEC_CH_UA_MOBILE,
|
||||
header.ACCEPT,
|
||||
header.ORIGIN,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.REFERER,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.COOKIE,
|
||||
header.PRIORITY,
|
||||
},
|
||||
}
|
||||
|
||||
var headerCache = profiles.NewHeaderCache(headerOrderDesktop, headerOrderMobile)
|
||||
|
||||
// --- Static header set (Variant.BuildHeaders) --------------------------------
|
||||
|
||||
// buildHeadersDesktop constructs the desktop Chrome 152 request header set.
|
||||
func buildHeadersDesktop(os profiles.OSKey) *g.MapOrd[g.String, g.String] {
|
||||
h := g.NewMapOrd[g.String, g.String]()
|
||||
h.Insert(":authority", "")
|
||||
h.Insert(":method", "")
|
||||
h.Insert(":path", "")
|
||||
h.Insert(":scheme", "")
|
||||
h.Insert(header.ACCEPT_ENCODING, "gzip, deflate, br, zstd")
|
||||
h.Insert(header.ACCEPT_LANGUAGE, "en-US,en;q=0.9")
|
||||
h.Insert(header.AUTHORIZATION, "")
|
||||
h.Insert(header.COOKIE, "")
|
||||
h.Insert(header.ORIGIN, "")
|
||||
h.Insert(header.REFERER, "")
|
||||
h.Insert(header.SEC_CH_UA, SecCHUA)
|
||||
h.Insert(header.SEC_CH_UA_MOBILE, os.Mobile())
|
||||
h.Insert(header.SEC_CH_UA_PLATFORM, Platform.Get(os).UnwrapOrDefault())
|
||||
h.Insert(header.USER_AGENT, UserAgent.Get(os).UnwrapOrDefault())
|
||||
|
||||
return &h
|
||||
}
|
||||
|
||||
// buildHeadersMobile constructs the placeholder mobile Chrome 152 request header set.
|
||||
// On the day real Chrome Android header set diverges from desktop (different Accept-Encoding,
|
||||
// shorter sec-ch-ua, different ordering / inserts), replace this body — it is the single point
|
||||
// of substitution for the entire mobile header set.
|
||||
func buildHeadersMobile(os profiles.OSKey) *g.MapOrd[g.String, g.String] {
|
||||
h := g.NewMapOrd[g.String, g.String]()
|
||||
h.Insert(":authority", "")
|
||||
h.Insert(":method", "")
|
||||
h.Insert(":path", "")
|
||||
h.Insert(":scheme", "")
|
||||
h.Insert(header.ACCEPT_ENCODING, "gzip, deflate, br, zstd")
|
||||
h.Insert(header.ACCEPT_LANGUAGE, "en-US,en;q=0.9")
|
||||
h.Insert(header.AUTHORIZATION, "")
|
||||
h.Insert(header.COOKIE, "")
|
||||
h.Insert(header.ORIGIN, "")
|
||||
h.Insert(header.REFERER, "")
|
||||
h.Insert(header.SEC_CH_UA, SecCHUA)
|
||||
h.Insert(header.SEC_CH_UA_MOBILE, os.Mobile())
|
||||
h.Insert(header.SEC_CH_UA_PLATFORM, Platform.Get(os).UnwrapOrDefault())
|
||||
h.Insert(header.USER_AGENT, UserAgent.Get(os).UnwrapOrDefault())
|
||||
|
||||
return &h
|
||||
}
|
||||
|
||||
// --- Per-request header pipeline (Variant.Headers) ---------------------------
|
||||
|
||||
// DesktopApplier applies the desktop Chrome request-header pipeline. Wired into chrome.Desktop.
|
||||
var DesktopApplier = profiles.NewApplier(insertDesktopHeaders, insertDesktopHeaders, headerCache, false)
|
||||
|
||||
// MobileApplier applies the mobile Chrome request-header pipeline. Wired into chrome.Mobile.
|
||||
var MobileApplier = profiles.NewApplier(insertMobileHeaders, insertMobileHeaders, headerCache, true)
|
||||
|
||||
func insertDesktopHeaders[T ~string](headers *g.MapOrd[T, T], method string) {
|
||||
switch method {
|
||||
case http.MethodPost:
|
||||
headers.Insert(header.ACCEPT, "*/*")
|
||||
headers.Insert(header.CACHE_CONTROL, "no-cache")
|
||||
headers.Insert(header.CONTENT_TYPE, "")
|
||||
headers.Insert(header.CONTENT_LENGTH, "")
|
||||
headers.Insert(header.PRAGMA, "no-cache")
|
||||
headers.Insert(header.PRIORITY, "u=1, i")
|
||||
headers.Insert(header.SEC_FETCH_DEST, "empty")
|
||||
headers.Insert(header.SEC_FETCH_MODE, "cors")
|
||||
headers.Insert(header.SEC_FETCH_SITE, "same-origin")
|
||||
default:
|
||||
headers.Insert(
|
||||
header.ACCEPT,
|
||||
"text/html,application/xhtml+xml,application/xml;q=0.9,image/avif,image/webp,image/apng,*/*;q=0.8,application/signed-exchange;v=b3;q=0.7",
|
||||
)
|
||||
headers.Insert(header.PRIORITY, "u=0, i")
|
||||
headers.Insert(header.SEC_FETCH_DEST, "document")
|
||||
headers.Insert(header.SEC_FETCH_MODE, "navigate")
|
||||
headers.Insert(header.SEC_FETCH_SITE, "none")
|
||||
headers.Insert(header.SEC_FETCH_USER, "?1")
|
||||
headers.Insert(header.UPGRADE_INSECURE_REQUESTS, "1")
|
||||
}
|
||||
}
|
||||
|
||||
// insertMobileHeaders is a placeholder mobile variant. On the day the real Chrome Android header
|
||||
// inserts diverge from desktop, this function is the single point to substitute them.
|
||||
func insertMobileHeaders[T ~string](headers *g.MapOrd[T, T], method string) {
|
||||
switch method {
|
||||
case http.MethodPost:
|
||||
headers.Insert(header.ACCEPT, "*/*")
|
||||
headers.Insert(header.CACHE_CONTROL, "no-cache")
|
||||
headers.Insert(header.CONTENT_TYPE, "")
|
||||
headers.Insert(header.CONTENT_LENGTH, "")
|
||||
headers.Insert(header.PRAGMA, "no-cache")
|
||||
headers.Insert(header.PRIORITY, "u=1, i")
|
||||
headers.Insert(header.SEC_FETCH_DEST, "empty")
|
||||
headers.Insert(header.SEC_FETCH_MODE, "cors")
|
||||
headers.Insert(header.SEC_FETCH_SITE, "same-origin")
|
||||
default:
|
||||
headers.Insert(
|
||||
header.ACCEPT,
|
||||
"text/html,application/xhtml+xml,application/xml;q=0.9,image/avif,image/webp,image/apng,*/*;q=0.8,application/signed-exchange;v=b3;q=0.7",
|
||||
)
|
||||
headers.Insert(header.PRIORITY, "u=0, i")
|
||||
headers.Insert(header.SEC_FETCH_DEST, "document")
|
||||
headers.Insert(header.SEC_FETCH_MODE, "navigate")
|
||||
headers.Insert(header.SEC_FETCH_SITE, "none")
|
||||
headers.Insert(header.SEC_FETCH_USER, "?1")
|
||||
headers.Insert(header.UPGRADE_INSECURE_REQUESTS, "1")
|
||||
}
|
||||
}
|
||||
+180
@@ -0,0 +1,180 @@
|
||||
package chrome
|
||||
|
||||
import utls "github.com/refraction-networking/utls"
|
||||
|
||||
// ML-DSA (post-quantum) TLS 1.3 signature schemes that Chrome 152 advertises at the
|
||||
// head of its signature_algorithms list. utls does not yet name them, so they are
|
||||
// spelled as raw SignatureScheme code points (per the TLS ML-DSA draft). Without
|
||||
// them a Chrome-152 UA ships a pre-152 JA4, which some Akamai deployments reject.
|
||||
const (
|
||||
MLDSA44 utls.SignatureScheme = 0x0904
|
||||
MLDSA65 utls.SignatureScheme = 0x0905
|
||||
MLDSA87 utls.SignatureScheme = 0x0906
|
||||
)
|
||||
|
||||
// extensionTrustAnchors is the TLS Trust Anchor Identifiers extension
|
||||
// (draft-ietf-tls-trust-anchor-ids, code point 0xca34). Chrome 152 sends it on every
|
||||
// ClientHello, so its presence is part of the Chrome 152 JA4 (t13d1517h2_…); a hello
|
||||
// without it is a pre-152 fingerprint, which some Akamai deployments reject when the
|
||||
// UA claims Chrome 152.
|
||||
const extensionTrustAnchors uint16 = 0xca34
|
||||
|
||||
// chrome152TrustAnchorIDs is the extension body Chrome 152 sends: a length-prefixed
|
||||
// TrustAnchorIdentifierList of the Chrome Root Store trust anchor IDs it advertises,
|
||||
// captured verbatim from a Google Chrome 152.0.7977 desktop ClientHello.
|
||||
//
|
||||
// The order matters and is not randomised. Chromium walks the compiled-in kChromeRootCertList
|
||||
// in array order (net/cert/internal/trust_store_chrome.cc), passes the bytes to
|
||||
// SSL_set1_requested_trust_anchors untouched, and BoringSSL writes them into the extension
|
||||
// verbatim — nothing shuffles the list the way the extension order is shuffled. So the order
|
||||
// is a compile-time constant of one Chrome build: stable across connections and hosts, but two
|
||||
// builds carrying different root store snapshots emit different orders for the same set of IDs.
|
||||
// These bytes must therefore come from a Google Chrome release, matching the branding the
|
||||
// profile claims in sec-ch-ua; a Chromium or ungoogled-chromium build of the same major
|
||||
// version is not interchangeable here.
|
||||
var chrome152TrustAnchorIDs = []byte{
|
||||
0x00, 0xcc,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x06,
|
||||
0x08, 0x83, 0x9a, 0x64, 0x8c, 0x9b, 0x2d, 0x01, 0x07,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x0c,
|
||||
0x05, 0x82, 0xdf, 0x13, 0x02, 0x06,
|
||||
0x05, 0x82, 0xdf, 0x13, 0x02, 0x13,
|
||||
0x08, 0x83, 0x9a, 0x64, 0x8c, 0x9b, 0x2d, 0x01, 0x0d,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x01,
|
||||
0x05, 0x82, 0xdf, 0x13, 0x02, 0x0d,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x0d,
|
||||
0x05, 0x82, 0xdf, 0x13, 0x02, 0x0f,
|
||||
0x08, 0x83, 0x9a, 0x64, 0x8c, 0x9b, 0x2d, 0x01, 0x08,
|
||||
0x05, 0x82, 0xdf, 0x13, 0x02, 0x12,
|
||||
0x08, 0x83, 0x9a, 0x64, 0x8c, 0x9b, 0x2d, 0x01, 0x09,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x02,
|
||||
0x05, 0x82, 0xdf, 0x13, 0x02, 0x01,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x0e,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x09,
|
||||
0x08, 0x83, 0x9a, 0x64, 0x8c, 0x9b, 0x2d, 0x01, 0x0a,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x03,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x0f,
|
||||
0x08, 0x83, 0x9a, 0x64, 0x8c, 0x9b, 0x2d, 0x01, 0x0b,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x04,
|
||||
0x05, 0x82, 0xdf, 0x13, 0x02, 0x14,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x0a,
|
||||
0x08, 0x83, 0x9a, 0x64, 0x8c, 0x9b, 0x2d, 0x01, 0x13,
|
||||
0x08, 0x83, 0x9a, 0x64, 0x8c, 0x9b, 0x2d, 0x01, 0x12,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x07,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x08,
|
||||
0x08, 0x83, 0x9a, 0x64, 0x8c, 0x9b, 0x2d, 0x01, 0x0c,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x05,
|
||||
0x05, 0x82, 0xdf, 0x13, 0x02, 0x0e,
|
||||
0x04, 0xd6, 0x79, 0x09, 0x0b,
|
||||
}
|
||||
|
||||
// HelloChrome_152 mirrors HelloChrome_150 (ML-DSA signature schemes at the head of
|
||||
// signature_algorithms) and adds what Chrome 152 sends on top: a GREASE signature scheme
|
||||
// and the trust_anchors extension. Together with the per-connection extension shuffle
|
||||
// (Variant.ShuffleExtensions / JA.Chrome152) this reproduces a Chrome 152 desktop
|
||||
// ClientHello structurally, not just its JA4.
|
||||
//
|
||||
// The extension order declared below is Chrome's own order. It is never shuffled in place:
|
||||
// the shuffle runs per connection on a private clone at dial time, so this value stays a
|
||||
// stable reference point across runs.
|
||||
var HelloChrome_152 = utls.ClientHelloSpec{
|
||||
CipherSuites: []uint16{
|
||||
utls.GREASE_PLACEHOLDER,
|
||||
utls.TLS_AES_128_GCM_SHA256,
|
||||
utls.TLS_AES_256_GCM_SHA384,
|
||||
utls.TLS_CHACHA20_POLY1305_SHA256,
|
||||
utls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256,
|
||||
utls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256,
|
||||
utls.TLS_ECDHE_ECDSA_WITH_AES_256_GCM_SHA384,
|
||||
utls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384,
|
||||
utls.TLS_ECDHE_ECDSA_WITH_CHACHA20_POLY1305,
|
||||
utls.TLS_ECDHE_RSA_WITH_CHACHA20_POLY1305,
|
||||
utls.TLS_ECDHE_RSA_WITH_AES_128_CBC_SHA,
|
||||
utls.TLS_ECDHE_RSA_WITH_AES_256_CBC_SHA,
|
||||
utls.TLS_RSA_WITH_AES_128_GCM_SHA256,
|
||||
utls.TLS_RSA_WITH_AES_256_GCM_SHA384,
|
||||
utls.TLS_RSA_WITH_AES_128_CBC_SHA,
|
||||
utls.TLS_RSA_WITH_AES_256_CBC_SHA,
|
||||
},
|
||||
CompressionMethods: []byte{0x00},
|
||||
Extensions: []utls.TLSExtension{
|
||||
&utls.UtlsGREASEExtension{},
|
||||
&utls.SNIExtension{},
|
||||
&utls.ExtendedMasterSecretExtension{},
|
||||
&utls.RenegotiationInfoExtension{
|
||||
Renegotiation: utls.RenegotiateOnceAsClient,
|
||||
},
|
||||
&utls.SupportedCurvesExtension{
|
||||
Curves: []utls.CurveID{
|
||||
utls.GREASE_PLACEHOLDER,
|
||||
utls.X25519MLKEM768,
|
||||
utls.X25519,
|
||||
utls.CurveP256,
|
||||
utls.CurveP384,
|
||||
},
|
||||
},
|
||||
&utls.SupportedPointsExtension{
|
||||
SupportedPoints: []byte{0x00},
|
||||
},
|
||||
&utls.SessionTicketExtension{},
|
||||
&utls.ALPNExtension{
|
||||
AlpnProtocols: []string{"h2", "http/1.1"},
|
||||
},
|
||||
&utls.StatusRequestExtension{},
|
||||
&utls.SignatureAlgorithmsExtension{
|
||||
SupportedSignatureAlgorithms: []utls.SignatureScheme{
|
||||
// Chrome 152 GREASEs signature_algorithms; surf substitutes the
|
||||
// placeholder with a random GREASE value per connection (see JA.getSpec).
|
||||
utls.GREASE_PLACEHOLDER,
|
||||
MLDSA44,
|
||||
MLDSA65,
|
||||
MLDSA87,
|
||||
utls.ECDSAWithP256AndSHA256,
|
||||
utls.PSSWithSHA256,
|
||||
utls.PKCS1WithSHA256,
|
||||
utls.ECDSAWithP384AndSHA384,
|
||||
utls.PSSWithSHA384,
|
||||
utls.PKCS1WithSHA384,
|
||||
utls.PSSWithSHA512,
|
||||
utls.PKCS1WithSHA512,
|
||||
},
|
||||
},
|
||||
&utls.SCTExtension{},
|
||||
&utls.KeyShareExtension{
|
||||
KeyShares: []utls.KeyShare{
|
||||
{Group: utls.GREASE_PLACEHOLDER, Data: []byte{0}},
|
||||
{Group: utls.X25519MLKEM768},
|
||||
{Group: utls.X25519},
|
||||
},
|
||||
},
|
||||
&utls.PSKKeyExchangeModesExtension{
|
||||
Modes: []uint8{
|
||||
utls.PskModeDHE,
|
||||
},
|
||||
},
|
||||
&utls.SupportedVersionsExtension{
|
||||
Versions: []uint16{
|
||||
utls.GREASE_PLACEHOLDER,
|
||||
utls.VersionTLS13,
|
||||
utls.VersionTLS12,
|
||||
},
|
||||
},
|
||||
&utls.UtlsCompressCertExtension{
|
||||
Algorithms: []utls.CertCompressionAlgo{
|
||||
utls.CertCompressionBrotli,
|
||||
},
|
||||
},
|
||||
&utls.ApplicationSettingsExtensionNew{
|
||||
SupportedProtocols: []string{"h2"},
|
||||
},
|
||||
&utls.GenericExtension{Id: extensionTrustAnchors, Data: chrome152TrustAnchorIDs},
|
||||
utls.BoringGREASEECH(),
|
||||
&utls.UtlsGREASEExtension{},
|
||||
&utls.UtlsPreSharedKeyExtension{},
|
||||
},
|
||||
}
|
||||
|
||||
// HelloChrome_152_Mobile is a placeholder mobile variant. On the day real Chrome Android 152
|
||||
// ClientHello bytes are observed, replace this body — it is the single point of substitution
|
||||
// for the mobile TLS fingerprint.
|
||||
var HelloChrome_152_Mobile = HelloChrome_152
|
||||
+54
@@ -0,0 +1,54 @@
|
||||
package chrome
|
||||
|
||||
import (
|
||||
"github.com/enetx/http2"
|
||||
"github.com/enetx/surf/profiles"
|
||||
)
|
||||
|
||||
// configureH2Desktop applies the desktop Chrome 152 HTTP/2 SETTINGS chain.
|
||||
func configureH2Desktop(h profiles.H2Config) {
|
||||
h.HeaderTableSize(65536).
|
||||
EnablePush(0).
|
||||
InitialWindowSize(6291456).
|
||||
MaxHeaderListSize(262144).
|
||||
ConnectionFlow(15663105).
|
||||
PriorityParam(http2.PriorityParam{
|
||||
StreamDep: 0,
|
||||
Exclusive: true,
|
||||
Weight: 255,
|
||||
})
|
||||
}
|
||||
|
||||
// configureH2Mobile applies the placeholder mobile Chrome 152 HTTP/2 SETTINGS chain.
|
||||
// On the day real Chrome Android 152 H/2 settings are observed, replace this body.
|
||||
func configureH2Mobile(h profiles.H2Config) {
|
||||
h.HeaderTableSize(65536).
|
||||
EnablePush(0).
|
||||
InitialWindowSize(6291456).
|
||||
MaxHeaderListSize(262144).
|
||||
ConnectionFlow(15663105).
|
||||
PriorityParam(http2.PriorityParam{
|
||||
StreamDep: 0,
|
||||
Exclusive: true,
|
||||
Weight: 255,
|
||||
})
|
||||
}
|
||||
|
||||
// configureH3Desktop applies the desktop Chrome 152 HTTP/3 SETTINGS chain.
|
||||
func configureH3Desktop(h profiles.H3Config) {
|
||||
h.QpackMaxTableCapacity(65536).
|
||||
MaxFieldSectionSize(262144).
|
||||
QpackBlockedStreams(100).
|
||||
SettingsH3Datagram(1).
|
||||
Grease()
|
||||
}
|
||||
|
||||
// configureH3Mobile applies the placeholder mobile Chrome 152 HTTP/3 SETTINGS chain.
|
||||
// On the day real Chrome Android 152 H/3 settings are observed, replace this body.
|
||||
func configureH3Mobile(h profiles.H3Config) {
|
||||
h.QpackMaxTableCapacity(65536).
|
||||
MaxFieldSectionSize(262144).
|
||||
QpackBlockedStreams(100).
|
||||
SettingsH3Datagram(1).
|
||||
Grease()
|
||||
}
|
||||
+32
@@ -0,0 +1,32 @@
|
||||
package chrome
|
||||
|
||||
import (
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/surf/profiles"
|
||||
)
|
||||
|
||||
// SecCHUA is the static value of the sec-ch-ua header for Chrome 152 (Chromium brand list).
|
||||
// If the mobile sec-ch-ua diverges from desktop, introduce SecCHUAMobile here and wire it into
|
||||
// chrome.Mobile in variant.go.
|
||||
const SecCHUA = `"Chromium";v="152", "Not?A_Brand";v="24", "Google Chrome";v="152"`
|
||||
|
||||
// UserAgent maps every supported impersonated OS to its Chrome 152 User-Agent string.
|
||||
// It is shared between Desktop and Mobile variants — UA strings are an OS property, not a
|
||||
// form-factor property (Chrome on Android always identifies as mobile, Chrome on Windows
|
||||
// always as desktop, regardless of which fingerprint variant is dispatched).
|
||||
var UserAgent = g.Map[profiles.OSKey, g.String]{
|
||||
profiles.Windows: "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/152.0.0.0 Safari/537.36",
|
||||
profiles.MacOS: "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/152.0.0.0 Safari/537.36",
|
||||
profiles.Linux: "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/152.0.0.0 Safari/537.36",
|
||||
profiles.Android: "Mozilla/5.0 (Linux; Android 10; K) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/152.0.0.0 Mobile Safari/537.36",
|
||||
profiles.IOS: "Mozilla/5.0 (iPad; CPU OS 26_3_0 like Mac OS X) AppleWebKit/605.1.15 (KHTML, like Gecko) CriOS/152.0.7977.53 Mobile/15E148 Safari/604.1",
|
||||
}
|
||||
|
||||
// Platform maps every supported impersonated OS to its sec-ch-ua-platform header value.
|
||||
var Platform = g.Map[profiles.OSKey, g.String]{
|
||||
profiles.Windows: `"Windows"`,
|
||||
profiles.MacOS: `"macOS"`,
|
||||
profiles.Linux: `"Linux"`,
|
||||
profiles.Android: `"Android"`,
|
||||
profiles.IOS: `"iOS"`,
|
||||
}
|
||||
+27
@@ -0,0 +1,27 @@
|
||||
package chrome
|
||||
|
||||
import "github.com/enetx/surf/profiles"
|
||||
|
||||
// Desktop is the Chrome 152 desktop variant - current production fingerprint.
|
||||
var Desktop = profiles.Variant{
|
||||
HelloSpec: &HelloChrome_152,
|
||||
ShuffleExtensions: true,
|
||||
Boundary: Boundary,
|
||||
ConfigureH2: configureH2Desktop,
|
||||
ConfigureH3: configureH3Desktop,
|
||||
BuildHeaders: buildHeadersDesktop,
|
||||
Headers: DesktopApplier,
|
||||
}
|
||||
|
||||
// Mobile is a placeholder Chrome 152 mobile variant. On the day real Chrome Android 152 bytes
|
||||
// are observed, replace HelloChrome_152_Mobile, configureH2Mobile, configureH3Mobile and
|
||||
// buildHeadersMobile bodies — the variant fields below stay as-is.
|
||||
var Mobile = profiles.Variant{
|
||||
HelloSpec: &HelloChrome_152_Mobile,
|
||||
ShuffleExtensions: true,
|
||||
Boundary: Boundary,
|
||||
ConfigureH2: configureH2Mobile,
|
||||
ConfigureH3: configureH3Mobile,
|
||||
BuildHeaders: buildHeadersMobile,
|
||||
Headers: MobileApplier,
|
||||
}
|
||||
+43
@@ -0,0 +1,43 @@
|
||||
package firefox
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
|
||||
"github.com/enetx/g"
|
||||
)
|
||||
|
||||
// Firefox implementation: https://github.com/mozilla/gecko-dev/blob/master/dom/html/HTMLFormSubmission.cpp#L355
|
||||
func Boundary() g.String {
|
||||
// C++
|
||||
// mBoundary.AssignLiteral("----geckoformboundary");
|
||||
// mBoundary.AppendInt(mozilla::RandomUint64OrDie(), 16);
|
||||
// mBoundary.AppendInt(mozilla::RandomUint64OrDie(), 16);
|
||||
|
||||
// prefix := "----geckoformboundary"
|
||||
// var num1, num2 uint64
|
||||
// binary.Read(rand.Reader, binary.BigEndian, &num1)
|
||||
// binary.Read(rand.Reader, binary.BigEndian, &num2)
|
||||
// return g.Sprintf("%s%x%x", prefix, num1, num2)
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// C++
|
||||
// mBoundary.AssignLiteral("---------------------------");
|
||||
// mBoundary.AppendInt(static_cast<uint32_t>(mozilla::RandomUint64OrDie()));
|
||||
// mBoundary.AppendInt(static_cast<uint32_t>(mozilla::RandomUint64OrDie()));
|
||||
// mBoundary.AppendInt(static_cast<uint32_t>(mozilla::RandomUint64OrDie()));
|
||||
|
||||
prefix := g.String("---------------------------")
|
||||
|
||||
var builder g.Builder
|
||||
builder.WriteString(prefix)
|
||||
|
||||
for range 3 {
|
||||
var b [4]byte
|
||||
rand.Read(b[:])
|
||||
builder.WriteString(g.Int(binary.LittleEndian.Uint32(b[:])).String())
|
||||
}
|
||||
|
||||
return builder.String()
|
||||
}
|
||||
+290
@@ -0,0 +1,290 @@
|
||||
package firefox
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/surf/header"
|
||||
"github.com/enetx/surf/profiles"
|
||||
)
|
||||
|
||||
// --- Header order maps -------------------------------------------------------
|
||||
|
||||
var headerOrderDesktop = g.Map[string, g.Slice[string]]{
|
||||
http.MethodGet: {
|
||||
":method",
|
||||
":path",
|
||||
":authority",
|
||||
":scheme",
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.REFERER,
|
||||
header.AUTHORIZATION,
|
||||
header.COOKIE,
|
||||
header.UPGRADE_INSECURE_REQUESTS,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_USER,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodGet + "http3": {
|
||||
":method",
|
||||
":scheme",
|
||||
":authority",
|
||||
":path",
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.REFERER,
|
||||
header.AUTHORIZATION,
|
||||
header.COOKIE,
|
||||
header.UPGRADE_INSECURE_REQUESTS,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_USER,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodPost: {
|
||||
":method",
|
||||
":path",
|
||||
":authority",
|
||||
":scheme",
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.REFERER,
|
||||
header.CONTENT_TYPE,
|
||||
header.AUTHORIZATION,
|
||||
header.CONTENT_LENGTH,
|
||||
header.ORIGIN,
|
||||
header.COOKIE,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.PRIORITY,
|
||||
header.PRAGMA,
|
||||
header.CACHE_CONTROL,
|
||||
},
|
||||
|
||||
http.MethodPost + "http3": {
|
||||
":method",
|
||||
":scheme",
|
||||
":authority",
|
||||
":path",
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.REFERER,
|
||||
header.CONTENT_TYPE,
|
||||
header.AUTHORIZATION,
|
||||
header.CONTENT_LENGTH,
|
||||
header.ORIGIN,
|
||||
header.COOKIE,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.PRIORITY,
|
||||
header.PRAGMA,
|
||||
header.CACHE_CONTROL,
|
||||
},
|
||||
}
|
||||
|
||||
// headerOrderMobile is a placeholder mobile variant. On the day real Firefox Android header
|
||||
// ordering is observed, this map is the single point to substitute it without touching desktop.
|
||||
// The literal is a physical copy of headerOrderDesktop so the two maps can diverge independently.
|
||||
var headerOrderMobile = g.Map[string, g.Slice[string]]{
|
||||
http.MethodGet: {
|
||||
":method",
|
||||
":path",
|
||||
":authority",
|
||||
":scheme",
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.REFERER,
|
||||
header.AUTHORIZATION,
|
||||
header.COOKIE,
|
||||
header.UPGRADE_INSECURE_REQUESTS,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_USER,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodGet + "http3": {
|
||||
":method",
|
||||
":scheme",
|
||||
":authority",
|
||||
":path",
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.REFERER,
|
||||
header.AUTHORIZATION,
|
||||
header.COOKIE,
|
||||
header.UPGRADE_INSECURE_REQUESTS,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.SEC_FETCH_USER,
|
||||
header.PRIORITY,
|
||||
},
|
||||
|
||||
http.MethodPost: {
|
||||
":method",
|
||||
":path",
|
||||
":authority",
|
||||
":scheme",
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.REFERER,
|
||||
header.CONTENT_TYPE,
|
||||
header.AUTHORIZATION,
|
||||
header.CONTENT_LENGTH,
|
||||
header.ORIGIN,
|
||||
header.COOKIE,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.PRIORITY,
|
||||
header.PRAGMA,
|
||||
header.CACHE_CONTROL,
|
||||
},
|
||||
|
||||
http.MethodPost + "http3": {
|
||||
":method",
|
||||
":scheme",
|
||||
":authority",
|
||||
":path",
|
||||
header.USER_AGENT,
|
||||
header.ACCEPT,
|
||||
header.ACCEPT_LANGUAGE,
|
||||
header.ACCEPT_ENCODING,
|
||||
header.REFERER,
|
||||
header.CONTENT_TYPE,
|
||||
header.AUTHORIZATION,
|
||||
header.CONTENT_LENGTH,
|
||||
header.ORIGIN,
|
||||
header.COOKIE,
|
||||
header.SEC_FETCH_DEST,
|
||||
header.SEC_FETCH_MODE,
|
||||
header.SEC_FETCH_SITE,
|
||||
header.PRIORITY,
|
||||
header.PRAGMA,
|
||||
header.CACHE_CONTROL,
|
||||
},
|
||||
}
|
||||
|
||||
var headerCache = profiles.NewHeaderCache(headerOrderDesktop, headerOrderMobile)
|
||||
|
||||
// --- Static header set (Variant.BuildHeaders) --------------------------------
|
||||
|
||||
// buildHeadersDesktop constructs the desktop Firefox 148 request header set.
|
||||
// Firefox does not emit Client Hints UA-CH headers (sec-ch-ua / sec-ch-ua-mobile /
|
||||
// sec-ch-ua-platform).
|
||||
func buildHeadersDesktop(os profiles.OSKey) *g.MapOrd[g.String, g.String] {
|
||||
h := g.NewMapOrd[g.String, g.String]()
|
||||
h.Insert(":authority", "")
|
||||
h.Insert(":method", "")
|
||||
h.Insert(":path", "")
|
||||
h.Insert(":scheme", "")
|
||||
h.Insert(header.ACCEPT_ENCODING, "gzip, deflate, br, zstd")
|
||||
h.Insert(header.ACCEPT_LANGUAGE, "en-US,en;q=0.5")
|
||||
h.Insert(header.AUTHORIZATION, "")
|
||||
h.Insert(header.COOKIE, "")
|
||||
h.Insert(header.ORIGIN, "")
|
||||
h.Insert(header.REFERER, "")
|
||||
h.Insert(header.USER_AGENT, UserAgent.Get(os).UnwrapOrDefault())
|
||||
|
||||
return &h
|
||||
}
|
||||
|
||||
// buildHeadersMobile constructs the placeholder mobile Firefox 148 request header set.
|
||||
// On the day real Firefox Android header set diverges from desktop, replace this body — it is
|
||||
// the single point of substitution for the entire mobile header set.
|
||||
func buildHeadersMobile(os profiles.OSKey) *g.MapOrd[g.String, g.String] {
|
||||
h := g.NewMapOrd[g.String, g.String]()
|
||||
h.Insert(":authority", "")
|
||||
h.Insert(":method", "")
|
||||
h.Insert(":path", "")
|
||||
h.Insert(":scheme", "")
|
||||
h.Insert(header.ACCEPT_ENCODING, "gzip, deflate, br, zstd")
|
||||
h.Insert(header.ACCEPT_LANGUAGE, "en-US,en;q=0.5")
|
||||
h.Insert(header.AUTHORIZATION, "")
|
||||
h.Insert(header.COOKIE, "")
|
||||
h.Insert(header.ORIGIN, "")
|
||||
h.Insert(header.REFERER, "")
|
||||
h.Insert(header.USER_AGENT, UserAgent.Get(os).UnwrapOrDefault())
|
||||
|
||||
return &h
|
||||
}
|
||||
|
||||
// --- Per-request header pipeline (Variant.Headers) ---------------------------
|
||||
|
||||
// DesktopApplier applies the desktop Firefox request-header pipeline. Wired into firefox.Desktop.
|
||||
var DesktopApplier = profiles.NewApplier(insertDesktopHeaders, insertDesktopHeaders, headerCache, false)
|
||||
|
||||
// MobileApplier applies the mobile Firefox request-header pipeline. Wired into firefox.Mobile.
|
||||
var MobileApplier = profiles.NewApplier(insertMobileHeaders, insertMobileHeaders, headerCache, true)
|
||||
|
||||
func insertDesktopHeaders[T ~string](headers *g.MapOrd[T, T], method string) {
|
||||
switch method {
|
||||
case http.MethodPost:
|
||||
headers.Insert(header.ACCEPT, "*/*")
|
||||
headers.Insert(header.CACHE_CONTROL, "no-cache")
|
||||
headers.Insert(header.CONTENT_TYPE, "")
|
||||
headers.Insert(header.CONTENT_LENGTH, "")
|
||||
headers.Insert(header.PRAGMA, "no-cache")
|
||||
headers.Insert(header.PRIORITY, "u=1, i")
|
||||
headers.Insert(header.SEC_FETCH_DEST, "empty")
|
||||
headers.Insert(header.SEC_FETCH_MODE, "cors")
|
||||
headers.Insert(header.SEC_FETCH_SITE, "same-origin")
|
||||
default:
|
||||
headers.Insert(header.ACCEPT, "text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8")
|
||||
headers.Insert(header.PRIORITY, "u=0, i")
|
||||
headers.Insert(header.SEC_FETCH_DEST, "document")
|
||||
headers.Insert(header.SEC_FETCH_MODE, "navigate")
|
||||
headers.Insert(header.SEC_FETCH_SITE, "none")
|
||||
headers.Insert(header.SEC_FETCH_USER, "?1")
|
||||
headers.Insert(header.UPGRADE_INSECURE_REQUESTS, "1")
|
||||
}
|
||||
}
|
||||
|
||||
// insertMobileHeaders is a placeholder mobile variant. On the day the real Firefox Android header
|
||||
// inserts diverge from desktop, this function is the single point to substitute them.
|
||||
func insertMobileHeaders[T ~string](headers *g.MapOrd[T, T], method string) {
|
||||
switch method {
|
||||
case http.MethodPost:
|
||||
headers.Insert(header.ACCEPT, "*/*")
|
||||
headers.Insert(header.CACHE_CONTROL, "no-cache")
|
||||
headers.Insert(header.CONTENT_TYPE, "")
|
||||
headers.Insert(header.CONTENT_LENGTH, "")
|
||||
headers.Insert(header.PRAGMA, "no-cache")
|
||||
headers.Insert(header.PRIORITY, "u=1, i")
|
||||
headers.Insert(header.SEC_FETCH_DEST, "empty")
|
||||
headers.Insert(header.SEC_FETCH_MODE, "cors")
|
||||
headers.Insert(header.SEC_FETCH_SITE, "same-origin")
|
||||
default:
|
||||
headers.Insert(header.ACCEPT, "text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8")
|
||||
headers.Insert(header.PRIORITY, "u=0, i")
|
||||
headers.Insert(header.SEC_FETCH_DEST, "document")
|
||||
headers.Insert(header.SEC_FETCH_MODE, "navigate")
|
||||
headers.Insert(header.SEC_FETCH_SITE, "none")
|
||||
headers.Insert(header.SEC_FETCH_USER, "?1")
|
||||
headers.Insert(header.UPGRADE_INSECURE_REQUESTS, "1")
|
||||
}
|
||||
}
|
||||
+10
@@ -0,0 +1,10 @@
|
||||
package firefox
|
||||
|
||||
import utls "github.com/refraction-networking/utls"
|
||||
|
||||
// HelloFirefox_148 is the desktop ClientHelloID for Firefox 148.
|
||||
var HelloFirefox_148 = utls.HelloFirefox_148
|
||||
|
||||
// HelloFirefox_148_Mobile is a placeholder mobile variant; functionally identical to the desktop
|
||||
// ClientHelloID until real Firefox Android bytes are observed and substituted here.
|
||||
var HelloFirefox_148_Mobile = utls.HelloFirefox_148
|
||||
+58
@@ -0,0 +1,58 @@
|
||||
package firefox
|
||||
|
||||
import (
|
||||
"github.com/enetx/http2"
|
||||
"github.com/enetx/surf/profiles"
|
||||
)
|
||||
|
||||
// configureH2Desktop applies the desktop Firefox 148 HTTP/2 SETTINGS chain.
|
||||
func configureH2Desktop(h profiles.H2Config) {
|
||||
h.InitialStreamID(3).
|
||||
HeaderTableSize(65536).
|
||||
EnablePush(0).
|
||||
InitialWindowSize(131072).
|
||||
MaxFrameSize(16384).
|
||||
ConnectionFlow(12517377).
|
||||
PriorityParam(http2.PriorityParam{
|
||||
StreamDep: 0,
|
||||
Exclusive: false,
|
||||
Weight: 41,
|
||||
})
|
||||
}
|
||||
|
||||
// configureH2Mobile applies the placeholder mobile Firefox 148 HTTP/2 SETTINGS chain.
|
||||
// On the day real Firefox Android 148 H/2 settings are observed, replace this body.
|
||||
func configureH2Mobile(h profiles.H2Config) {
|
||||
h.InitialStreamID(3).
|
||||
HeaderTableSize(65536).
|
||||
EnablePush(0).
|
||||
InitialWindowSize(131072).
|
||||
MaxFrameSize(16384).
|
||||
ConnectionFlow(12517377).
|
||||
PriorityParam(http2.PriorityParam{
|
||||
StreamDep: 0,
|
||||
Exclusive: false,
|
||||
Weight: 41,
|
||||
})
|
||||
}
|
||||
|
||||
// configureH3Desktop applies the desktop Firefox 148 HTTP/3 SETTINGS chain.
|
||||
func configureH3Desktop(h profiles.H3Config) {
|
||||
h.QpackMaxTableCapacity(65536).
|
||||
QpackBlockedStreams(20).
|
||||
EnableWebtransport(0).
|
||||
H3Datagram(1).
|
||||
SettingsH3Datagram(1).
|
||||
EnableConnectProtocol(1)
|
||||
}
|
||||
|
||||
// configureH3Mobile applies the placeholder mobile Firefox 148 HTTP/3 SETTINGS chain.
|
||||
// On the day real Firefox Android 148 H/3 settings are observed, replace this body.
|
||||
func configureH3Mobile(h profiles.H3Config) {
|
||||
h.QpackMaxTableCapacity(65536).
|
||||
QpackBlockedStreams(20).
|
||||
EnableWebtransport(0).
|
||||
H3Datagram(1).
|
||||
SettingsH3Datagram(1).
|
||||
EnableConnectProtocol(1)
|
||||
}
|
||||
+17
@@ -0,0 +1,17 @@
|
||||
package firefox
|
||||
|
||||
import (
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/surf/profiles"
|
||||
)
|
||||
|
||||
// UserAgent maps every supported impersonated OS to its Firefox 148 User-Agent string.
|
||||
// Shared between Desktop and Mobile variants — UA strings are an OS property, not a
|
||||
// form-factor property. iOS Firefox uses the FxiOS variant (under the hood it is WebKit).
|
||||
var UserAgent = g.Map[profiles.OSKey, g.String]{
|
||||
profiles.Windows: "Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:148.0) Gecko/20100101 Firefox/148.0",
|
||||
profiles.MacOS: "Mozilla/5.0 (Macintosh; Intel Mac OS X 10.15; rv:148.0) Gecko/20100101 Firefox/148.0",
|
||||
profiles.Linux: "Mozilla/5.0 (X11; Linux x86_64; rv:148.0) Gecko/20100101 Firefox/148.0",
|
||||
profiles.Android: "Mozilla/5.0 (Android 16; Mobile; rv:148.0) Gecko/148.0 Firefox/148.0",
|
||||
profiles.IOS: "Mozilla/5.0 (iPhone; CPU iPhone OS 18_7 like Mac OS X) AppleWebKit/605.1.15 (KHTML, like Gecko) FxiOS/148.0 Mobile/15E148 Safari/605.1.15",
|
||||
}
|
||||
+25
@@ -0,0 +1,25 @@
|
||||
package firefox
|
||||
|
||||
import "github.com/enetx/surf/profiles"
|
||||
|
||||
// Desktop is the Firefox 148 desktop variant — current production fingerprint.
|
||||
var Desktop = profiles.Variant{
|
||||
HelloID: HelloFirefox_148,
|
||||
Boundary: Boundary,
|
||||
ConfigureH2: configureH2Desktop,
|
||||
ConfigureH3: configureH3Desktop,
|
||||
BuildHeaders: buildHeadersDesktop,
|
||||
Headers: DesktopApplier,
|
||||
}
|
||||
|
||||
// Mobile is a placeholder Firefox 148 mobile variant. On the day real Firefox Android 148 bytes
|
||||
// are observed, replace HelloFirefox_148_Mobile, configureH2Mobile, configureH3Mobile and
|
||||
// buildHeadersMobile bodies — the variant fields below stay as-is.
|
||||
var Mobile = profiles.Variant{
|
||||
HelloID: HelloFirefox_148_Mobile,
|
||||
Boundary: Boundary,
|
||||
ConfigureH2: configureH2Mobile,
|
||||
ConfigureH3: configureH3Mobile,
|
||||
BuildHeaders: buildHeadersMobile,
|
||||
Headers: MobileApplier,
|
||||
}
|
||||
+100
@@ -0,0 +1,100 @@
|
||||
package profiles
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/g/cmp"
|
||||
)
|
||||
|
||||
// HeadersApplier applies the browser-specific request-header pipeline (insert defaults +
|
||||
// reorder by header-order map) to a request header map. Form-factor (desktop/mobile) is baked
|
||||
// into each applier instance — Variant.Desktop and Variant.Mobile carry different appliers.
|
||||
type HeadersApplier func(headers any, method string)
|
||||
|
||||
// HeadersFn is the type of a generic headers function instantiated for a concrete T ~string.
|
||||
type HeadersFn[T ~string] func(*g.MapOrd[T, T], string)
|
||||
|
||||
// NewApplier wraps a browser-specific header-insert pair with the standard "insert, then
|
||||
// reorder by per-method enum" pipeline shared by every profile. The form factor (desktop/mobile)
|
||||
// is baked in via cache.Enums(mobile); the lookup runs on each call so HeaderCache lazy-init is
|
||||
// preserved.
|
||||
//
|
||||
// The constructor takes the two concrete HeadersFn instantiations as separate parameters so
|
||||
// callers can pass the same generic function twice and let Go infer T at each site.
|
||||
func NewApplier(
|
||||
insertG HeadersFn[g.String],
|
||||
insertS HeadersFn[string],
|
||||
cache *HeaderCache,
|
||||
mobile bool,
|
||||
) HeadersApplier {
|
||||
return func(h any, method string) {
|
||||
switch v := h.(type) {
|
||||
case *g.MapOrd[g.String, g.String]:
|
||||
insertG(v, method)
|
||||
SortByOrder(v, method, cache.Enums(mobile))
|
||||
case *g.MapOrd[string, string]:
|
||||
insertS(v, method)
|
||||
SortByOrder(v, method, cache.Enums(mobile))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// HeaderCache lazily builds and caches per-method header-position enums for both desktop and
|
||||
// mobile header-order maps. Each browser profile constructs one HeaderCache from its own
|
||||
// headerOrderDesktop / headerOrderMobile literals. The dispatcher asks for enums via Enums(mobile).
|
||||
type HeaderCache struct {
|
||||
desktopOrder g.Map[string, g.Slice[string]]
|
||||
mobileOrder g.Map[string, g.Slice[string]]
|
||||
|
||||
desktopEnums g.Map[string, g.MapOrd[string, g.Int]]
|
||||
mobileEnums g.Map[string, g.MapOrd[string, g.Int]]
|
||||
|
||||
onceDesktop sync.Once
|
||||
onceMobile sync.Once
|
||||
}
|
||||
|
||||
// NewHeaderCache wires the cache to two header-order maps. Enums are built on first access.
|
||||
func NewHeaderCache(desktop, mobile g.Map[string, g.Slice[string]]) *HeaderCache {
|
||||
return &HeaderCache{desktopOrder: desktop, mobileOrder: mobile}
|
||||
}
|
||||
|
||||
// Enums returns the desktop or mobile enum map (lazy-built on first call).
|
||||
func (c *HeaderCache) Enums(mobile bool) g.Map[string, g.MapOrd[string, g.Int]] {
|
||||
if mobile {
|
||||
c.onceMobile.Do(func() { c.mobileEnums = buildHeaderEnums(c.mobileOrder) })
|
||||
return c.mobileEnums
|
||||
}
|
||||
|
||||
c.onceDesktop.Do(func() { c.desktopEnums = buildHeaderEnums(c.desktopOrder) })
|
||||
return c.desktopEnums
|
||||
}
|
||||
|
||||
func buildHeaderEnums(order g.Map[string, g.Slice[string]]) g.Map[string, g.MapOrd[string, g.Int]] {
|
||||
enums := g.NewMap[string, g.MapOrd[string, g.Int]]()
|
||||
|
||||
for method, headers := range order {
|
||||
h := g.NewMapOrd[string, g.Int]()
|
||||
headers.Iter().Enumerate().
|
||||
ForEach(func(k g.Int, v string) {
|
||||
h.Insert(v, k)
|
||||
})
|
||||
|
||||
enums[method] = h
|
||||
}
|
||||
|
||||
return enums
|
||||
}
|
||||
|
||||
// SortByOrder reorders headers in-place according to the per-method enum positions returned
|
||||
// by HeaderCache.Enums. Profile packages call it from headersDesktop / headersMobile after the
|
||||
// browser-specific Insert step. Falls back to the GET enum when method is not present.
|
||||
func SortByOrder[T ~string](h *g.MapOrd[T, T], method string, enums g.Map[string, g.MapOrd[string, g.Int]]) {
|
||||
enum := enums.Get(method).UnwrapOr(enums["GET"])
|
||||
|
||||
h.SortByKey(func(a, b T) cmp.Ordering {
|
||||
ida := enum.Get(string(a))
|
||||
idb := enum.Get(string(b))
|
||||
return ida.UnwrapOrDefault().Cmp(idb.UnwrapOrDefault())
|
||||
})
|
||||
}
|
||||
+102
@@ -0,0 +1,102 @@
|
||||
// Package profiles defines the shared variant contract used by browser profile packages
|
||||
// (profiles/chrome, profiles/firefox, ...) and consumed by the top-level surf package.
|
||||
//
|
||||
// A Variant is the self-contained description of one browser and form-factor combination:
|
||||
// TLS ClientHello, HTTP/2 + HTTP/3 SETTINGS, header set, User-Agent / sec-ch-ua data per OS.
|
||||
// The surf.Impersonate dispatcher reads chrome.Desktop / chrome.Mobile / firefox.Desktop /
|
||||
// firefox.Mobile values, applies the static fields, and runs ConfigureH2/ConfigureH3 against
|
||||
// adapters that satisfy H2Config / H3Config — without the profile package importing surf.
|
||||
package profiles
|
||||
|
||||
import (
|
||||
"github.com/enetx/g"
|
||||
"github.com/enetx/http2"
|
||||
utls "github.com/refraction-networking/utls"
|
||||
)
|
||||
|
||||
// OSKey identifies the impersonated operating system. Surf's Impersonate stores it directly
|
||||
// (no separate ImpersonateOS enum). Profile packages use it as a lookup key for UA / Platform.
|
||||
type OSKey int
|
||||
|
||||
const (
|
||||
Windows OSKey = iota
|
||||
MacOS
|
||||
Linux
|
||||
Android
|
||||
IOS
|
||||
)
|
||||
|
||||
// IsMobile reports whether the OS is a mobile form factor (Android or iOS).
|
||||
func (k OSKey) IsMobile() bool { return k == Android || k == IOS }
|
||||
|
||||
// Mobile returns the value of the sec-ch-ua-mobile header: "?1" for mobile OS, "?0" otherwise.
|
||||
func (k OSKey) Mobile() g.String {
|
||||
if k.IsMobile() {
|
||||
return "?1"
|
||||
}
|
||||
|
||||
return "?0"
|
||||
}
|
||||
|
||||
// H2Config is the fluent contract used by profile.ConfigureH2 callbacks. It mirrors the
|
||||
// methods on surf.HTTP2Settings, the surf package provides an adapter that satisfies it.
|
||||
type H2Config interface {
|
||||
HeaderTableSize(uint32) H2Config
|
||||
EnablePush(uint32) H2Config
|
||||
MaxConcurrentStreams(uint32) H2Config
|
||||
InitialWindowSize(uint32) H2Config
|
||||
MaxFrameSize(uint32) H2Config
|
||||
MaxHeaderListSize(uint32) H2Config
|
||||
NoRFC7540Priorities(uint32) H2Config
|
||||
ConnectionFlow(uint32) H2Config
|
||||
InitialStreamID(uint32) H2Config
|
||||
PriorityParam(http2.PriorityParam) H2Config
|
||||
PriorityFrames([]http2.PriorityFrame) H2Config
|
||||
}
|
||||
|
||||
// H3Config is the fluent contract used by profile.ConfigureH3 callbacks. It mirrors the
|
||||
// methods on surf.HTTP3Settings, the surf package provides an adapter that satisfies it.
|
||||
type H3Config interface {
|
||||
QpackMaxTableCapacity(uint64) H3Config
|
||||
MaxFieldSectionSize(uint64) H3Config
|
||||
QpackBlockedStreams(uint64) H3Config
|
||||
EnableConnectProtocol(uint64) H3Config
|
||||
SettingsH3Datagram(uint64) H3Config
|
||||
H3Datagram(uint64) H3Config
|
||||
EnableWebtransport(uint64) H3Config
|
||||
Grease() H3Config
|
||||
}
|
||||
|
||||
// Variant is a self-contained description of one browser and form-factor combination.
|
||||
type Variant struct {
|
||||
// HelloSpec takes precedence over HelloID when non-nil.
|
||||
HelloSpec *utls.ClientHelloSpec
|
||||
HelloID utls.ClientHelloID
|
||||
|
||||
// ShuffleExtensions re-shuffles HelloSpec's extension order on every connection, as
|
||||
// Chromium does since v110. Firefox keeps a fixed order.
|
||||
ShuffleExtensions bool
|
||||
|
||||
// Boundary is the multipart boundary generator for this browser. Same reference for both
|
||||
// Desktop and Mobile within one profile package (boundary is a per-browser property).
|
||||
Boundary func() g.String
|
||||
|
||||
// ConfigureH2 / ConfigureH3 own the fluent SETTINGS chain as code. The surf dispatcher
|
||||
// invokes them with adapters and calls Set() afterwards.
|
||||
ConfigureH2 func(H2Config)
|
||||
ConfigureH3 func(H3Config)
|
||||
|
||||
// BuildHeaders constructs the full ordered header map for one request — pseudo-headers,
|
||||
// Accept-Encoding, Accept-Language, authorization/cookie/origin/referer placeholders,
|
||||
// sec-ch-ua-* (Chromium only), User-Agent. Profile packages provide one BuildHeaders per
|
||||
// Variant so each browser and form-factor has its own single point of substitution for the
|
||||
// entire header set (set, values, and order are all browser-specific).
|
||||
BuildHeaders func(OSKey) *g.MapOrd[g.String, g.String]
|
||||
|
||||
// Headers applies the per-request header-order pipeline (the same Headers[T ~string]
|
||||
// function profile packages export) for the surf request path. Same value for Desktop and
|
||||
// Mobile within one profile package — header pipeline differs by mobile bool, not by
|
||||
// Variant — but living on Variant lets surf.Builder route through one indirection without
|
||||
// importing concrete profile packages from the request hot path.
|
||||
Headers HeadersApplier
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
debug
|
||||
debug.test
|
||||
main
|
||||
mockgen_tmp.go
|
||||
*.qtr
|
||||
*.qlog
|
||||
*.sqlog
|
||||
*.txt
|
||||
race.[0-9]*
|
||||
|
||||
fuzzing/*/*.zip
|
||||
fuzzing/*/coverprofile
|
||||
fuzzing/*/crashers
|
||||
fuzzing/*/sonarprofile
|
||||
fuzzing/*/suppressions
|
||||
fuzzing/*/corpus/
|
||||
|
||||
**/testdata/fuzz/
|
||||
|
||||
gomock_reflect_*/
|
||||
+99
@@ -0,0 +1,99 @@
|
||||
version: "2"
|
||||
linters:
|
||||
default: none
|
||||
enable:
|
||||
- asciicheck
|
||||
- copyloopvar
|
||||
- depguard
|
||||
- exhaustive
|
||||
- govet
|
||||
- ineffassign
|
||||
- misspell
|
||||
- nolintlint
|
||||
- prealloc
|
||||
- staticcheck
|
||||
- unconvert
|
||||
- unparam
|
||||
- unused
|
||||
- usetesting
|
||||
settings:
|
||||
depguard:
|
||||
rules:
|
||||
random:
|
||||
deny:
|
||||
- pkg: "math/rand$"
|
||||
desc: use math/rand/v2
|
||||
- pkg: "golang.org/x/exp/rand"
|
||||
desc: use math/rand/v2
|
||||
quicvarint:
|
||||
list-mode: strict
|
||||
files:
|
||||
- '**/github.com/quic-go/quic-go/quicvarint/*'
|
||||
- '!$test'
|
||||
allow:
|
||||
- $gostd
|
||||
rsa:
|
||||
list-mode: original
|
||||
deny:
|
||||
- pkg: crypto/rsa
|
||||
desc: "use crypto/ed25519 instead"
|
||||
ginkgo:
|
||||
list-mode: original
|
||||
deny:
|
||||
- pkg: github.com/onsi/ginkgo
|
||||
desc: "use standard Go tests"
|
||||
- pkg: github.com/onsi/ginkgo/v2
|
||||
desc: "use standard Go tests"
|
||||
- pkg: github.com/onsi/gomega
|
||||
desc: "use standard Go tests"
|
||||
http3-internal:
|
||||
list-mode: lax
|
||||
files:
|
||||
- '**/http3/**'
|
||||
deny:
|
||||
- pkg: 'github.com/quic-go/quic-go/internal'
|
||||
desc: 'no dependency on quic-go/internal'
|
||||
misspell:
|
||||
ignore-rules:
|
||||
- ect
|
||||
# see https://github.com/ldez/usetesting/issues/10
|
||||
usetesting:
|
||||
context-background: false
|
||||
context-todo: false
|
||||
exclusions:
|
||||
generated: lax
|
||||
presets:
|
||||
- comments
|
||||
- common-false-positives
|
||||
- legacy
|
||||
- std-error-handling
|
||||
rules:
|
||||
- linters:
|
||||
- depguard
|
||||
path: internal/qtls
|
||||
- linters:
|
||||
- exhaustive
|
||||
- prealloc
|
||||
- unparam
|
||||
path: _test\.go
|
||||
- linters:
|
||||
- staticcheck
|
||||
path: _test\.go
|
||||
text: 'SA1029:' # inappropriate key in call to context.WithValue
|
||||
paths:
|
||||
- internal/handshake/cipher_suite.go
|
||||
- third_party$
|
||||
- builtin$
|
||||
- examples$
|
||||
formatters:
|
||||
enable:
|
||||
- gofmt
|
||||
- gofumpt
|
||||
- goimports
|
||||
exclusions:
|
||||
generated: lax
|
||||
paths:
|
||||
- internal/handshake/cipher_suite.go
|
||||
- third_party$
|
||||
- builtin$
|
||||
- examples$
|
||||
+37
@@ -0,0 +1,37 @@
|
||||
# FIPS 140-3
|
||||
|
||||
quic-go relies on the Go standard library for cryptography, including the Go Cryptographic Module described in [The FIPS 140-3 Go Cryptographic Module](https://go.dev/blog/fips140). quic-go does not seek separate FIPS 140-3 validation as a cryptographic module. This document explains how quic-go uses Go standard library cryptography for QUIC operations relevant to FIPS 140-3.
|
||||
|
||||
Starting with quic-go v0.60, the behavior described here applies when built with Go 1.26 or newer. With older Go versions, quic-go still builds and runs as usual, without any attempt to meet FIPS 140 requirements.
|
||||
|
||||
## QUIC operations relevant to FIPS 140-3
|
||||
|
||||
quic-go delegates the TLS 1.3 handshake, certificate handling, cipher suite selection, session tickets, and the TLS key schedule to `crypto/tls`. When Go's FIPS 140-3 mode is active, `crypto/tls` restricts the algorithms it negotiates.
|
||||
|
||||
### Packet protection AEADs
|
||||
|
||||
The main quic-go-specific FIPS-relevant operations are the AEADs protecting Handshake, 0-RTT, and 1-RTT packets.
|
||||
|
||||
AES-GCM packet protection AEADs are constructed through the Go standard library's TLS 1.3 AES-GCM implementation. Today this uses `go:linkname` to call the unexported `crypto/tls.aeadAESGCMTLS13`, because the standard library does not yet expose a QUIC-specific constructor; see [golang/go#79219](https://github.com/golang/go/issues/79219).
|
||||
|
||||
ChaCha20-Poly1305 is not used in Go's FIPS 140-3 mode. `crypto/tls` avoids that cipher suite during negotiation, and quic-go additionally guards its internal ChaCha20-Poly1305 path when FIPS 140-3 mode is enabled.
|
||||
|
||||
### Header protection
|
||||
|
||||
For Handshake, 0-RTT, and 1-RTT packets protected with AES cipher suites, header protection keys are derived with `crypto/hkdf` and the AES block operation uses `crypto/aes`. ChaCha20 header protection is tied to the ChaCha20-Poly1305 cipher suite and is not reachable in FIPS 140-3 mode.
|
||||
|
||||
### Address validation tokens
|
||||
|
||||
quic-go encrypts the address validation tokens it sends in Retry packets and NEW_TOKEN frames. These are not TLS session tickets (those are handled by `crypto/tls`); they carry server-defined state such as the client address, timestamp, RTT information, and Retry connection IDs.
|
||||
|
||||
Token-protection keys are derived with `crypto/hkdf`, AES is used via `crypto/aes`, and the token AEAD is constructed with `cipher.NewGCMWithRandomNonce`, keeping token encryption on standard library primitives.
|
||||
|
||||
## QUIC operations not relevant to FIPS 140-3
|
||||
|
||||
### Initial packet protection
|
||||
|
||||
Initial packet protection (including Initial header protection) is not treated as FIPS 140-relevant confidentiality protection: the Initial secrets are derived from constants in RFC 9001 and the packet's destination connection ID, so any observer can derive the same keys. quic-go therefore disables strict FIPS 140 enforcement around Initial packet construction in Go 1.26 FIPS 140-3 mode. See the IETF QUIC mailing list discussion at <https://mailarchive.ietf.org/arch/msg/quic/k2kl2W_n5WDEZBbt3O31Ef2XBbM/>.
|
||||
|
||||
### Retry packet integrity tag
|
||||
|
||||
RFC 9001 defines the Retry packet integrity tag using fixed keys and nonces. It guards against accidental corruption and casual injection but does not encrypt packet contents. quic-go treats it as outside the FIPS 140 scope and disables strict FIPS 140 enforcement for that AEAD construction in Go 1.26 FIPS 140-3 mode.
|
||||
+59
@@ -0,0 +1,59 @@
|
||||
# Fuzzing
|
||||
|
||||
[](https://introspector.oss-fuzz.com/project-profile?project=quic-go)
|
||||
[](https://app.codecov.io/gh/quic-go/quic-go?flags%5B0%5D=clusterfuzz)
|
||||
[](https://app.codecov.io/gh/quic-go/quic-go?flags%5B0%5D=clusterfuzz-lite-batch)
|
||||
|
||||
Run the commands below from a local [`google/oss-fuzz`](https://github.com/google/oss-fuzz) checkout.
|
||||
Fuzz target names match the binary names listed in `oss-fuzz.sh` (for example, `frame_fuzzer_v2`).
|
||||
|
||||
Update the base images:
|
||||
```sh
|
||||
python3 infra/helper.py pull_images
|
||||
```
|
||||
|
||||
## Running fuzzers locally
|
||||
|
||||
The following steps run a single fuzz target and then open its line-by-line coverage in `go tool cover`.
|
||||
|
||||
```sh
|
||||
export DOCKER_DEFAULT_PLATFORM=linux/amd64
|
||||
export FUZZ_TARGET=<fuzz_target>
|
||||
export CORPUS_DIR=corpus/$FUZZ_TARGET
|
||||
|
||||
mkdir -p "$CORPUS_DIR"
|
||||
|
||||
python3 infra/helper.py build_image --no-pull quic-go
|
||||
python3 infra/helper.py build_fuzzers --sanitizer address quic-go
|
||||
python3 infra/helper.py run_fuzzer --corpus-dir="$CORPUS_DIR" quic-go "$FUZZ_TARGET"
|
||||
```
|
||||
|
||||
Leave `run_fuzzer` running for a while to build up a corpus. It unpacks the seed corpus zip into the corpus directory and appends new entries as it discovers them.
|
||||
|
||||
```sh
|
||||
python3 infra/helper.py build_fuzzers --sanitizer coverage quic-go
|
||||
python3 infra/helper.py coverage --no-serve --fuzz-target "$FUZZ_TARGET" --corpus-dir="$CORPUS_DIR" quic-go
|
||||
sed "s#^/out/#$(pwd)/build/out/quic-go/#" build/out/quic-go/fuzz.cov > "/tmp/quic-go-$FUZZ_TARGET.coverprofile"
|
||||
go tool cover -html="/tmp/quic-go-$FUZZ_TARGET.coverprofile"
|
||||
```
|
||||
|
||||
The `sed` command rewrites the container paths in `fuzz.cov` so that `go tool cover` can locate the source files in the local checkout.
|
||||
|
||||
To produce a coverage report against a modified local source tree, mount the local checkout when building the coverage fuzzers, the same way you would for reproducers:
|
||||
|
||||
```sh
|
||||
python3 infra/helper.py build_fuzzers --sanitizer coverage --mount_path /root/go/src/github.com/quic-go/quic-go quic-go <local_quic_go_dir>
|
||||
```
|
||||
|
||||
## Reproducing an OSS-Fuzz testcase
|
||||
|
||||
Download the reproducer file from the OSS-Fuzz report. To test a local fix, rebuild the fuzzers with the modified quic-go checkout mounted at the path expected by `oss-fuzz.sh`:
|
||||
|
||||
```sh
|
||||
export DOCKER_DEFAULT_PLATFORM=linux/amd64
|
||||
export FUZZ_TARGET=<fuzz_target>
|
||||
|
||||
python3 infra/helper.py build_image --no-pull quic-go
|
||||
python3 infra/helper.py build_fuzzers --sanitizer address --mount_path /root/go/src/github.com/quic-go/quic-go quic-go <local_quic_go_dir>
|
||||
python3 infra/helper.py reproduce quic-go "$FUZZ_TARGET" <reproducer_file>
|
||||
```
|
||||
+21
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2016 the quic-go authors & Google, Inc.
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
+65
@@ -0,0 +1,65 @@
|
||||
<div align="center" style="margin-bottom: 15px;">
|
||||
<img src="./assets/quic-go-logo.png" width="700" height="auto">
|
||||
</div>
|
||||
|
||||
# A QUIC implementation in pure Go
|
||||
|
||||
|
||||
[](https://quic-go.net/docs/)
|
||||
[](https://pkg.go.dev/github.com/quic-go/quic-go)
|
||||
[](https://codecov.io/gh/quic-go/quic-go/)
|
||||
[](https://issues.oss-fuzz.com/issues?q=quic-go)
|
||||
|
||||
quic-go is an implementation of the QUIC protocol ([RFC 9000](https://datatracker.ietf.org/doc/html/rfc9000), [RFC 9001](https://datatracker.ietf.org/doc/html/rfc9001), [RFC 9002](https://datatracker.ietf.org/doc/html/rfc9002)) in Go. It has support for HTTP/3 ([RFC 9114](https://datatracker.ietf.org/doc/html/rfc9114)), including QPACK ([RFC 9204](https://datatracker.ietf.org/doc/html/rfc9204)) and HTTP Datagrams ([RFC 9297](https://datatracker.ietf.org/doc/html/rfc9297)).
|
||||
|
||||
In addition to these base RFCs, it also implements the following RFCs:
|
||||
|
||||
* Unreliable Datagram Extension ([RFC 9221](https://datatracker.ietf.org/doc/html/rfc9221))
|
||||
* Datagram Packetization Layer Path MTU Discovery (DPLPMTUD, [RFC 8899](https://datatracker.ietf.org/doc/html/rfc8899))
|
||||
* QUIC Version 2 ([RFC 9369](https://datatracker.ietf.org/doc/html/rfc9369))
|
||||
* QUIC Event Logging using qlog ([draft-ietf-quic-qlog-main-schema](https://datatracker.ietf.org/doc/draft-ietf-quic-qlog-main-schema/) and [draft-ietf-quic-qlog-quic-events](https://datatracker.ietf.org/doc/draft-ietf-quic-qlog-quic-events/))
|
||||
* QUIC Stream Resets with Partial Delivery ([draft-ietf-quic-reliable-stream-reset-07](https://datatracker.ietf.org/doc/html/draft-ietf-quic-reliable-stream-reset-07) and [draft-ietf-quic-reliable-stream-reset-09](https://datatracker.ietf.org/doc/html/draft-ietf-quic-reliable-stream-reset-09))
|
||||
|
||||
Support for WebTransport over HTTP/3 ([draft-ietf-webtrans-http3](https://datatracker.ietf.org/doc/draft-ietf-webtrans-http3/)) is implemented in [webtransport-go](https://github.com/quic-go/webtransport-go).
|
||||
|
||||
Detailed documentation can be found on [quic-go.net](https://quic-go.net/docs/).
|
||||
|
||||
## FIPS 140-3
|
||||
|
||||
Starting with v0.60, quic-go supports use in FIPS 140-3 environments when built with Go 1.26 or newer, using Go standard library cryptography for the QUIC code paths relevant in FIPS mode; see [FIPS140.md](FIPS140.md) for details.
|
||||
|
||||
## Projects using quic-go
|
||||
|
||||
| Project | Description | Stars |
|
||||
| ---------------------------------------------------------- | --------------------------------------------------------------------------------------------------------------------------------------------------------------------- | --------------------------------------------------------------------------------------------------- |
|
||||
| [AdGuardHome](https://github.com/AdguardTeam/AdGuardHome) | Free and open source, powerful network-wide ads & trackers blocking DNS server. |  |
|
||||
| [algernon](https://github.com/xyproto/algernon) | Small self-contained pure-Go web server with Lua, Markdown, HTTP/2, QUIC, Redis and PostgreSQL support |  |
|
||||
| [caddy](https://github.com/caddyserver/caddy/) | Fast, multi-platform web server with automatic HTTPS |  |
|
||||
| [cloudflared](https://github.com/cloudflare/cloudflared) | A tunneling daemon that proxies traffic from the Cloudflare network to your origins |  |
|
||||
| [frp](https://github.com/fatedier/frp) | A fast reverse proxy to help you expose a local server behind a NAT or firewall to the internet |  |
|
||||
| [go-libp2p](https://github.com/libp2p/go-libp2p) | libp2p implementation in Go, powering [Kubo](https://github.com/ipfs/kubo) (IPFS) and [Lotus](https://github.com/filecoin-project/lotus) (Filecoin), among others |  |
|
||||
| [gost](https://github.com/go-gost/gost) | A simple security tunnel written in Go |  |
|
||||
| [Hysteria](https://github.com/apernet/hysteria) | A powerful, lightning fast and censorship resistant proxy |  |
|
||||
| [Mercure](https://github.com/dunglas/mercure) | An open, easy, fast, reliable and battery-efficient solution for real-time communications |  |
|
||||
| [nodepass](https://github.com/NodePassProject/nodepass) | A secure, efficient TCP/UDP tunneling solution that delivers fast, reliable access across network restrictions using pre-established TCP/QUIC/WebSocket or HTTP/2 connections. |  |
|
||||
| [OONI Probe](https://github.com/ooni/probe-cli) | Next generation OONI Probe. Library and CLI tool. |  |
|
||||
| [reverst](https://github.com/flipt-io/reverst) | Reverse Tunnels in Go over HTTP/3 and QUIC |  |
|
||||
| [RoadRunner](https://github.com/roadrunner-server/roadrunner) | High-performance PHP application server, process manager written in Go and powered with plugins |  |
|
||||
| [syncthing](https://github.com/syncthing/syncthing/) | Open Source Continuous File Synchronization |  |
|
||||
| [traefik](https://github.com/traefik/traefik) | The Cloud Native Application Proxy |  |
|
||||
| [v2ray-core](https://github.com/v2fly/v2ray-core) | A platform for building proxies to bypass network restrictions |  |
|
||||
| [YoMo](https://github.com/yomorun/yomo) | Streaming Serverless Framework for Geo-distributed System |  |
|
||||
|
||||
If you'd like to see your project added to this list, please send us a PR.
|
||||
|
||||
## Release Policy
|
||||
|
||||
quic-go always aims to support the latest two Go releases.
|
||||
|
||||
## Contributing
|
||||
|
||||
We are always happy to welcome new contributors! We have a number of self-contained issues that are suitable for first-time contributors, they are tagged with [help wanted](https://github.com/quic-go/quic-go/issues?q=is%3Aissue+is%3Aopen+label%3A%22help+wanted%22). If you have any questions, please feel free to reach out by opening an issue or leaving a comment.
|
||||
|
||||
## License
|
||||
|
||||
The code is licensed under the MIT license. The logo and brand assets are excluded from the MIT license. See [assets/LICENSE.md](https://github.com/quic-go/quic-go/tree/master/assets/LICENSE.md) for the full usage policy and details.
|
||||
+14
@@ -0,0 +1,14 @@
|
||||
# Security Policy
|
||||
|
||||
quic-go is an implementation of the QUIC protocol and related standards. No software is perfect, and we take reports of potential security issues very seriously.
|
||||
|
||||
## Reporting a Vulnerability
|
||||
|
||||
If you discover a vulnerability that could affect production deployments (e.g., a remotely exploitable issue), please report it [**privately**](https://github.com/quic-go/quic-go/security/advisories/new).
|
||||
Please **DO NOT file a public issue** for exploitable vulnerabilities.
|
||||
|
||||
If the issue is theoretical, non-exploitable, or related to an experimental feature, you may discuss it openly by filing a regular issue.
|
||||
|
||||
## Reporting a non-security bug
|
||||
|
||||
For bugs, feature requests, or other non-security concerns, please open a GitHub [issue](https://github.com/quic-go/quic-go/issues/new).
|
||||
+92
@@ -0,0 +1,92 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
)
|
||||
|
||||
type packetBuffer struct {
|
||||
Data []byte
|
||||
|
||||
// refCount counts how many packets Data is used in.
|
||||
// It doesn't support concurrent use.
|
||||
// It is > 1 when used for coalesced packet.
|
||||
refCount int
|
||||
}
|
||||
|
||||
// Split increases the refCount.
|
||||
// It must be called when a packet buffer is used for more than one packet,
|
||||
// e.g. when splitting coalesced packets.
|
||||
func (b *packetBuffer) Split() {
|
||||
b.refCount++
|
||||
}
|
||||
|
||||
// Decrement decrements the reference counter.
|
||||
// It doesn't put the buffer back into the pool.
|
||||
func (b *packetBuffer) Decrement() {
|
||||
b.refCount--
|
||||
if b.refCount < 0 {
|
||||
panic("negative packetBuffer refCount")
|
||||
}
|
||||
}
|
||||
|
||||
// MaybeRelease puts the packet buffer back into the pool,
|
||||
// if the reference counter already reached 0.
|
||||
func (b *packetBuffer) MaybeRelease() {
|
||||
// only put the packetBuffer back if it's not used any more
|
||||
if b.refCount == 0 {
|
||||
b.putBack()
|
||||
}
|
||||
}
|
||||
|
||||
// Release puts back the packet buffer into the pool.
|
||||
// It should be called when processing is definitely finished.
|
||||
func (b *packetBuffer) Release() {
|
||||
b.Decrement()
|
||||
if b.refCount != 0 {
|
||||
panic("packetBuffer refCount not zero")
|
||||
}
|
||||
b.putBack()
|
||||
}
|
||||
|
||||
// Len returns the length of Data
|
||||
func (b *packetBuffer) Len() protocol.ByteCount { return protocol.ByteCount(len(b.Data)) }
|
||||
func (b *packetBuffer) Cap() protocol.ByteCount { return protocol.ByteCount(cap(b.Data)) }
|
||||
|
||||
func (b *packetBuffer) putBack() {
|
||||
if cap(b.Data) == protocol.MaxPacketBufferSize {
|
||||
bufferPool.Put(b)
|
||||
return
|
||||
}
|
||||
if cap(b.Data) == protocol.MaxLargePacketBufferSize {
|
||||
largeBufferPool.Put(b)
|
||||
return
|
||||
}
|
||||
panic("putPacketBuffer called with packet of wrong size!")
|
||||
}
|
||||
|
||||
var bufferPool, largeBufferPool sync.Pool
|
||||
|
||||
func getPacketBuffer() *packetBuffer {
|
||||
buf := bufferPool.Get().(*packetBuffer)
|
||||
buf.refCount = 1
|
||||
buf.Data = buf.Data[:0]
|
||||
return buf
|
||||
}
|
||||
|
||||
func getLargePacketBuffer() *packetBuffer {
|
||||
buf := largeBufferPool.Get().(*packetBuffer)
|
||||
buf.refCount = 1
|
||||
buf.Data = buf.Data[:0]
|
||||
return buf
|
||||
}
|
||||
|
||||
func init() {
|
||||
bufferPool.New = func() any {
|
||||
return &packetBuffer{Data: make([]byte, 0, protocol.MaxPacketBufferSize)}
|
||||
}
|
||||
largeBufferPool.New = func() any {
|
||||
return &packetBuffer{Data: make([]byte, 0, protocol.MaxLargePacketBufferSize)}
|
||||
}
|
||||
}
|
||||
+109
@@ -0,0 +1,109 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"net"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
)
|
||||
|
||||
// make it possible to mock connection ID for initial generation in the tests
|
||||
var generateConnectionIDForInitial = protocol.GenerateConnectionIDForInitial
|
||||
|
||||
// DialAddr establishes a new QUIC connection to a server.
|
||||
// It resolves the address, and then creates a new UDP connection to dial the QUIC server.
|
||||
// When the QUIC connection is closed, this UDP connection is closed.
|
||||
// See [Dial] for more details.
|
||||
func DialAddr(ctx context.Context, addr string, tlsConf *tls.Config, conf *Config) (*Conn, error) {
|
||||
udpConn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4zero, Port: 0})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
udpAddr, err := net.ResolveUDPAddr("udp", addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tr, err := setupTransport(udpConn, tlsConf, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
conn, err := tr.dial(ctx, udpAddr, addr, tlsConf, conf, false)
|
||||
if err != nil {
|
||||
tr.Close()
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// DialAddrEarly establishes a new 0-RTT QUIC connection to a server.
|
||||
// See [DialAddr] for more details.
|
||||
func DialAddrEarly(ctx context.Context, addr string, tlsConf *tls.Config, conf *Config) (*Conn, error) {
|
||||
udpConn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4zero, Port: 0})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
udpAddr, err := net.ResolveUDPAddr("udp", addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tr, err := setupTransport(udpConn, tlsConf, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
conn, err := tr.dial(ctx, udpAddr, addr, tlsConf, conf, true)
|
||||
if err != nil {
|
||||
tr.Close()
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// DialEarly establishes a new 0-RTT QUIC connection to a server using a net.PacketConn.
|
||||
// See [Dial] for more details.
|
||||
func DialEarly(ctx context.Context, c net.PacketConn, addr net.Addr, tlsConf *tls.Config, conf *Config) (*Conn, error) {
|
||||
dl, err := setupTransport(c, tlsConf, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
conn, err := dl.DialEarly(ctx, addr, tlsConf, conf)
|
||||
if err != nil {
|
||||
dl.Close()
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// Dial establishes a new QUIC connection to a server using a net.PacketConn.
|
||||
// If the PacketConn satisfies the [OOBCapablePacketConn] interface (as a [net.UDPConn] does),
|
||||
// ECN and packet info support will be enabled. In this case, ReadMsgUDP and WriteMsgUDP
|
||||
// will be used instead of ReadFrom and WriteTo to read/write packets.
|
||||
// The [tls.Config] must define an application protocol (using tls.Config.NextProtos).
|
||||
//
|
||||
// This is a convenience function. More advanced use cases should instantiate a [Transport],
|
||||
// which offers configuration options for a more fine-grained control of the connection establishment,
|
||||
// including reusing the underlying UDP socket for multiple QUIC connections.
|
||||
func Dial(ctx context.Context, c net.PacketConn, addr net.Addr, tlsConf *tls.Config, conf *Config) (*Conn, error) {
|
||||
dl, err := setupTransport(c, tlsConf, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
conn, err := dl.Dial(ctx, addr, tlsConf, conf)
|
||||
if err != nil {
|
||||
dl.Close()
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func setupTransport(c net.PacketConn, tlsConf *tls.Config, createdPacketConn bool) (*Transport, error) {
|
||||
if tlsConf == nil {
|
||||
return nil, errors.New("quic: tls.Config not set")
|
||||
}
|
||||
return &Transport{
|
||||
Conn: c,
|
||||
createdConn: createdPacketConn,
|
||||
isSingleUse: true,
|
||||
}, nil
|
||||
}
|
||||
+58
@@ -0,0 +1,58 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"math/bits"
|
||||
"net"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/utils"
|
||||
)
|
||||
|
||||
// A closedLocalConn is a connection that we closed locally.
|
||||
// When receiving packets for such a connection, we need to retransmit the packet containing the CONNECTION_CLOSE frame,
|
||||
// with an exponential backoff.
|
||||
type closedLocalConn struct {
|
||||
counter atomic.Uint32
|
||||
logger utils.Logger
|
||||
|
||||
sendPacket func(net.Addr, packetInfo)
|
||||
}
|
||||
|
||||
var _ packetHandler = &closedLocalConn{}
|
||||
|
||||
// newClosedLocalConn creates a new closedLocalConn and runs it.
|
||||
func newClosedLocalConn(sendPacket func(net.Addr, packetInfo), logger utils.Logger) packetHandler {
|
||||
return &closedLocalConn{
|
||||
sendPacket: sendPacket,
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *closedLocalConn) handlePacket(p receivedPacket) {
|
||||
n := c.counter.Add(1)
|
||||
// exponential backoff
|
||||
// only send a CONNECTION_CLOSE for the 1st, 2nd, 4th, 8th, 16th, ... packet arriving
|
||||
if bits.OnesCount32(n) != 1 {
|
||||
return
|
||||
}
|
||||
c.logger.Debugf("Received %d packets after sending CONNECTION_CLOSE. Retransmitting.", n)
|
||||
c.sendPacket(p.remoteAddr, p.info)
|
||||
}
|
||||
|
||||
func (c *closedLocalConn) destroy(error) {}
|
||||
func (c *closedLocalConn) closeWithTransportError(TransportErrorCode) {}
|
||||
|
||||
// A closedRemoteConn is a connection that was closed remotely.
|
||||
// For such a connection, we might receive reordered packets that were sent before the CONNECTION_CLOSE.
|
||||
// We can just ignore those packets.
|
||||
type closedRemoteConn struct{}
|
||||
|
||||
var _ packetHandler = &closedRemoteConn{}
|
||||
|
||||
func newClosedRemoteConn() packetHandler {
|
||||
return &closedRemoteConn{}
|
||||
}
|
||||
|
||||
func (c *closedRemoteConn) handlePacket(receivedPacket) {}
|
||||
func (c *closedRemoteConn) destroy(error) {}
|
||||
func (c *closedRemoteConn) closeWithTransportError(TransportErrorCode) {}
|
||||
+23
@@ -0,0 +1,23 @@
|
||||
coverage:
|
||||
round: nearest
|
||||
ignore:
|
||||
- http3/gzip_reader.go
|
||||
- example/
|
||||
- interop/
|
||||
- internal/handshake/cipher_suite.go
|
||||
- internal/mocks/
|
||||
- internal/utils/linkedlist/linkedlist.go
|
||||
- internal/testdata
|
||||
- testutils/
|
||||
- fuzzing/
|
||||
- metrics/
|
||||
status:
|
||||
project:
|
||||
default:
|
||||
threshold: 0.5
|
||||
patch: false
|
||||
flags:
|
||||
clusterfuzz-lite-batch:
|
||||
joined: false
|
||||
clusterfuzz:
|
||||
joined: false
|
||||
+129
@@ -0,0 +1,129 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/quicvarint"
|
||||
)
|
||||
|
||||
// Clone clones a Config.
|
||||
func (c *Config) Clone() *Config {
|
||||
copy := *c
|
||||
return ©
|
||||
}
|
||||
|
||||
func (c *Config) handshakeTimeout() time.Duration {
|
||||
return 2 * c.HandshakeIdleTimeout
|
||||
}
|
||||
|
||||
func (c *Config) maxRetryTokenAge() time.Duration {
|
||||
return c.handshakeTimeout()
|
||||
}
|
||||
|
||||
func validateConfig(config *Config) error {
|
||||
if config == nil {
|
||||
return nil
|
||||
}
|
||||
const maxStreams = 1 << 60
|
||||
if config.MaxIncomingStreams > maxStreams {
|
||||
config.MaxIncomingStreams = maxStreams
|
||||
}
|
||||
if config.MaxIncomingUniStreams > maxStreams {
|
||||
config.MaxIncomingUniStreams = maxStreams
|
||||
}
|
||||
if config.MaxStreamReceiveWindow > quicvarint.Max {
|
||||
config.MaxStreamReceiveWindow = quicvarint.Max
|
||||
}
|
||||
if config.MaxConnectionReceiveWindow > quicvarint.Max {
|
||||
config.MaxConnectionReceiveWindow = quicvarint.Max
|
||||
}
|
||||
if config.InitialPacketSize > 0 && config.InitialPacketSize < protocol.MinInitialPacketSize {
|
||||
config.InitialPacketSize = protocol.MinInitialPacketSize
|
||||
}
|
||||
if config.InitialPacketSize > protocol.MaxPacketBufferSize {
|
||||
config.InitialPacketSize = protocol.MaxPacketBufferSize
|
||||
}
|
||||
// check that all QUIC versions are actually supported
|
||||
for _, v := range config.Versions {
|
||||
if !protocol.IsValidVersion(v) {
|
||||
return fmt.Errorf("invalid QUIC version: %s", v)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// populateConfig populates fields in the quic.Config with their default values, if none are set
|
||||
// it may be called with nil
|
||||
func populateConfig(config *Config) *Config {
|
||||
if config == nil {
|
||||
config = &Config{}
|
||||
}
|
||||
versions := config.Versions
|
||||
if len(versions) == 0 {
|
||||
versions = protocol.SupportedVersions
|
||||
}
|
||||
handshakeIdleTimeout := protocol.DefaultHandshakeIdleTimeout
|
||||
if config.HandshakeIdleTimeout != 0 {
|
||||
handshakeIdleTimeout = config.HandshakeIdleTimeout
|
||||
}
|
||||
idleTimeout := protocol.DefaultIdleTimeout
|
||||
if config.MaxIdleTimeout != 0 {
|
||||
idleTimeout = config.MaxIdleTimeout
|
||||
}
|
||||
initialStreamReceiveWindow := config.InitialStreamReceiveWindow
|
||||
if initialStreamReceiveWindow == 0 {
|
||||
initialStreamReceiveWindow = protocol.DefaultInitialMaxStreamData
|
||||
}
|
||||
maxStreamReceiveWindow := config.MaxStreamReceiveWindow
|
||||
if maxStreamReceiveWindow == 0 {
|
||||
maxStreamReceiveWindow = protocol.DefaultMaxReceiveStreamFlowControlWindow
|
||||
}
|
||||
initialConnectionReceiveWindow := config.InitialConnectionReceiveWindow
|
||||
if initialConnectionReceiveWindow == 0 {
|
||||
initialConnectionReceiveWindow = protocol.DefaultInitialMaxData
|
||||
}
|
||||
maxConnectionReceiveWindow := config.MaxConnectionReceiveWindow
|
||||
if maxConnectionReceiveWindow == 0 {
|
||||
maxConnectionReceiveWindow = protocol.DefaultMaxReceiveConnectionFlowControlWindow
|
||||
}
|
||||
maxIncomingStreams := config.MaxIncomingStreams
|
||||
if maxIncomingStreams == 0 {
|
||||
maxIncomingStreams = protocol.DefaultMaxIncomingStreams
|
||||
} else if maxIncomingStreams < 0 {
|
||||
maxIncomingStreams = 0
|
||||
}
|
||||
maxIncomingUniStreams := config.MaxIncomingUniStreams
|
||||
if maxIncomingUniStreams == 0 {
|
||||
maxIncomingUniStreams = protocol.DefaultMaxIncomingUniStreams
|
||||
} else if maxIncomingUniStreams < 0 {
|
||||
maxIncomingUniStreams = 0
|
||||
}
|
||||
initialPacketSize := config.InitialPacketSize
|
||||
if initialPacketSize == 0 {
|
||||
initialPacketSize = protocol.InitialPacketSize
|
||||
}
|
||||
|
||||
return &Config{
|
||||
GetConfigForClient: config.GetConfigForClient,
|
||||
Versions: versions,
|
||||
HandshakeIdleTimeout: handshakeIdleTimeout,
|
||||
MaxIdleTimeout: idleTimeout,
|
||||
KeepAlivePeriod: config.KeepAlivePeriod,
|
||||
InitialStreamReceiveWindow: initialStreamReceiveWindow,
|
||||
MaxStreamReceiveWindow: maxStreamReceiveWindow,
|
||||
InitialConnectionReceiveWindow: initialConnectionReceiveWindow,
|
||||
MaxConnectionReceiveWindow: maxConnectionReceiveWindow,
|
||||
AllowConnectionWindowIncrease: config.AllowConnectionWindowIncrease,
|
||||
MaxIncomingStreams: maxIncomingStreams,
|
||||
MaxIncomingUniStreams: maxIncomingUniStreams,
|
||||
TokenStore: config.TokenStore,
|
||||
EnableDatagrams: config.EnableDatagrams,
|
||||
InitialPacketSize: initialPacketSize,
|
||||
DisablePathMTUDiscovery: config.DisablePathMTUDiscovery,
|
||||
EnableStreamResetPartialDelivery: config.EnableStreamResetPartialDelivery,
|
||||
Allow0RTT: config.Allow0RTT,
|
||||
Tracer: config.Tracer,
|
||||
}
|
||||
}
|
||||
+212
@@ -0,0 +1,212 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"slices"
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/monotime"
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/internal/qerr"
|
||||
"github.com/quic-go/quic-go/internal/wire"
|
||||
)
|
||||
|
||||
type connRunnerCallbacks struct {
|
||||
AddConnectionID func(protocol.ConnectionID)
|
||||
RemoveConnectionID func(protocol.ConnectionID)
|
||||
ReplaceWithClosed func([]protocol.ConnectionID, []byte, time.Duration)
|
||||
}
|
||||
|
||||
// The memory address of the Transport is used as the key.
|
||||
type connRunners map[connRunner]connRunnerCallbacks
|
||||
|
||||
func (cr connRunners) AddConnectionID(id protocol.ConnectionID) {
|
||||
for _, c := range cr {
|
||||
c.AddConnectionID(id)
|
||||
}
|
||||
}
|
||||
|
||||
func (cr connRunners) RemoveConnectionID(id protocol.ConnectionID) {
|
||||
for _, c := range cr {
|
||||
c.RemoveConnectionID(id)
|
||||
}
|
||||
}
|
||||
|
||||
func (cr connRunners) ReplaceWithClosed(ids []protocol.ConnectionID, b []byte, expiry time.Duration) {
|
||||
for _, c := range cr {
|
||||
c.ReplaceWithClosed(ids, b, expiry)
|
||||
}
|
||||
}
|
||||
|
||||
type connIDToRetire struct {
|
||||
t monotime.Time
|
||||
connID protocol.ConnectionID
|
||||
}
|
||||
|
||||
type connIDGenerator struct {
|
||||
generator ConnectionIDGenerator
|
||||
highestSeq uint64
|
||||
connRunners connRunners
|
||||
|
||||
activeSrcConnIDs map[uint64]protocol.ConnectionID
|
||||
connIDsToRetire []connIDToRetire // sorted by t
|
||||
initialClientDestConnID *protocol.ConnectionID // nil for the client
|
||||
|
||||
statelessResetter *statelessResetter
|
||||
|
||||
queueControlFrame func(wire.Frame)
|
||||
}
|
||||
|
||||
func newConnIDGenerator(
|
||||
runner connRunner,
|
||||
initialConnectionID protocol.ConnectionID,
|
||||
initialClientDestConnID *protocol.ConnectionID, // nil for the client
|
||||
statelessResetter *statelessResetter,
|
||||
callbacks connRunnerCallbacks,
|
||||
queueControlFrame func(wire.Frame),
|
||||
generator ConnectionIDGenerator,
|
||||
) *connIDGenerator {
|
||||
m := &connIDGenerator{
|
||||
generator: generator,
|
||||
activeSrcConnIDs: make(map[uint64]protocol.ConnectionID),
|
||||
statelessResetter: statelessResetter,
|
||||
connRunners: map[connRunner]connRunnerCallbacks{runner: callbacks},
|
||||
queueControlFrame: queueControlFrame,
|
||||
}
|
||||
m.activeSrcConnIDs[0] = initialConnectionID
|
||||
m.initialClientDestConnID = initialClientDestConnID
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *connIDGenerator) SetMaxActiveConnIDs(limit uint64) error {
|
||||
if m.generator.ConnectionIDLen() == 0 {
|
||||
return nil
|
||||
}
|
||||
// The active_connection_id_limit transport parameter is the number of
|
||||
// connection IDs the peer will store. This limit includes the connection ID
|
||||
// used during the handshake, and the one sent in the preferred_address
|
||||
// transport parameter.
|
||||
// We currently don't send the preferred_address transport parameter,
|
||||
// so we can issue (limit - 1) connection IDs.
|
||||
for i := uint64(len(m.activeSrcConnIDs)); i < min(limit, protocol.MaxIssuedConnectionIDs); i++ {
|
||||
if err := m.issueNewConnID(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *connIDGenerator) Retire(seq uint64, sentWithDestConnID protocol.ConnectionID, expiry monotime.Time) error {
|
||||
if seq > m.highestSeq {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.ProtocolViolation,
|
||||
ErrorMessage: fmt.Sprintf("retired connection ID %d (highest issued: %d)", seq, m.highestSeq),
|
||||
}
|
||||
}
|
||||
connID, ok := m.activeSrcConnIDs[seq]
|
||||
// We might already have deleted this connection ID, if this is a duplicate frame.
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
if connID == sentWithDestConnID {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.ProtocolViolation,
|
||||
ErrorMessage: fmt.Sprintf("retired connection ID %d (%s), which was used as the Destination Connection ID on this packet", seq, connID),
|
||||
}
|
||||
}
|
||||
m.queueConnIDForRetiring(connID, expiry)
|
||||
|
||||
delete(m.activeSrcConnIDs, seq)
|
||||
// Don't issue a replacement for the initial connection ID.
|
||||
if seq == 0 {
|
||||
return nil
|
||||
}
|
||||
return m.issueNewConnID()
|
||||
}
|
||||
|
||||
func (m *connIDGenerator) queueConnIDForRetiring(connID protocol.ConnectionID, expiry monotime.Time) {
|
||||
idx := slices.IndexFunc(m.connIDsToRetire, func(c connIDToRetire) bool {
|
||||
return c.t.After(expiry)
|
||||
})
|
||||
if idx == -1 {
|
||||
idx = len(m.connIDsToRetire)
|
||||
}
|
||||
m.connIDsToRetire = slices.Insert(m.connIDsToRetire, idx, connIDToRetire{t: expiry, connID: connID})
|
||||
}
|
||||
|
||||
func (m *connIDGenerator) issueNewConnID() error {
|
||||
connID, err := m.generator.GenerateConnectionID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
m.activeSrcConnIDs[m.highestSeq+1] = connID
|
||||
m.connRunners.AddConnectionID(connID)
|
||||
m.queueControlFrame(&wire.NewConnectionIDFrame{
|
||||
SequenceNumber: m.highestSeq + 1,
|
||||
ConnectionID: connID,
|
||||
StatelessResetToken: m.statelessResetter.GetStatelessResetToken(connID),
|
||||
})
|
||||
m.highestSeq++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *connIDGenerator) SetHandshakeComplete(connIDExpiry monotime.Time) {
|
||||
if m.initialClientDestConnID != nil {
|
||||
m.queueConnIDForRetiring(*m.initialClientDestConnID, connIDExpiry)
|
||||
m.initialClientDestConnID = nil
|
||||
}
|
||||
}
|
||||
|
||||
func (m *connIDGenerator) RemoveRetiredConnIDs(now monotime.Time) {
|
||||
if len(m.connIDsToRetire) == 0 {
|
||||
return
|
||||
}
|
||||
for _, c := range m.connIDsToRetire {
|
||||
if c.t.After(now) {
|
||||
break
|
||||
}
|
||||
m.connRunners.RemoveConnectionID(c.connID)
|
||||
m.connIDsToRetire = m.connIDsToRetire[1:]
|
||||
}
|
||||
}
|
||||
|
||||
func (m *connIDGenerator) RemoveAll() {
|
||||
if m.initialClientDestConnID != nil {
|
||||
m.connRunners.RemoveConnectionID(*m.initialClientDestConnID)
|
||||
}
|
||||
for _, connID := range m.activeSrcConnIDs {
|
||||
m.connRunners.RemoveConnectionID(connID)
|
||||
}
|
||||
for _, c := range m.connIDsToRetire {
|
||||
m.connRunners.RemoveConnectionID(c.connID)
|
||||
}
|
||||
}
|
||||
|
||||
func (m *connIDGenerator) ReplaceWithClosed(connClose []byte, expiry time.Duration) {
|
||||
connIDs := make([]protocol.ConnectionID, 0, len(m.activeSrcConnIDs)+len(m.connIDsToRetire)+1)
|
||||
if m.initialClientDestConnID != nil {
|
||||
connIDs = append(connIDs, *m.initialClientDestConnID)
|
||||
}
|
||||
for _, connID := range m.activeSrcConnIDs {
|
||||
connIDs = append(connIDs, connID)
|
||||
}
|
||||
for _, c := range m.connIDsToRetire {
|
||||
connIDs = append(connIDs, c.connID)
|
||||
}
|
||||
m.connRunners.ReplaceWithClosed(connIDs, connClose, expiry)
|
||||
}
|
||||
|
||||
func (m *connIDGenerator) AddConnRunner(runner connRunner, r connRunnerCallbacks) {
|
||||
// The transport might have already been added earlier.
|
||||
// This happens if the application migrates back to and old path.
|
||||
if _, ok := m.connRunners[runner]; ok {
|
||||
return
|
||||
}
|
||||
m.connRunners[runner] = r
|
||||
if m.initialClientDestConnID != nil {
|
||||
r.AddConnectionID(*m.initialClientDestConnID)
|
||||
}
|
||||
for _, connID := range m.activeSrcConnIDs {
|
||||
r.AddConnectionID(connID)
|
||||
}
|
||||
}
|
||||
+321
@@ -0,0 +1,321 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"slices"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/internal/qerr"
|
||||
"github.com/quic-go/quic-go/internal/utils"
|
||||
"github.com/quic-go/quic-go/internal/wire"
|
||||
)
|
||||
|
||||
type newConnID struct {
|
||||
SequenceNumber uint64
|
||||
ConnectionID protocol.ConnectionID
|
||||
StatelessResetToken protocol.StatelessResetToken
|
||||
}
|
||||
|
||||
type connIDManager struct {
|
||||
queue []newConnID
|
||||
|
||||
highestProbingID uint64
|
||||
pathProbing map[pathID]newConnID // initialized lazily
|
||||
|
||||
handshakeComplete bool
|
||||
activeSequenceNumber uint64
|
||||
highestRetired uint64
|
||||
activeConnectionID protocol.ConnectionID
|
||||
activeStatelessResetToken *protocol.StatelessResetToken
|
||||
|
||||
// We change the connection ID after sending on average
|
||||
// protocol.PacketsPerConnectionID packets. The actual value is randomized
|
||||
// hide the packet loss rate from on-path observers.
|
||||
rand utils.Rand
|
||||
packetsSinceLastChange uint32
|
||||
packetsPerConnectionID uint32
|
||||
|
||||
addStatelessResetToken func(protocol.StatelessResetToken)
|
||||
removeStatelessResetToken func(protocol.StatelessResetToken)
|
||||
queueControlFrame func(wire.Frame)
|
||||
|
||||
closed bool
|
||||
}
|
||||
|
||||
func newConnIDManager(
|
||||
initialDestConnID protocol.ConnectionID,
|
||||
addStatelessResetToken func(protocol.StatelessResetToken),
|
||||
removeStatelessResetToken func(protocol.StatelessResetToken),
|
||||
queueControlFrame func(wire.Frame),
|
||||
) *connIDManager {
|
||||
return &connIDManager{
|
||||
activeConnectionID: initialDestConnID,
|
||||
addStatelessResetToken: addStatelessResetToken,
|
||||
removeStatelessResetToken: removeStatelessResetToken,
|
||||
queueControlFrame: queueControlFrame,
|
||||
queue: make([]newConnID, 0, protocol.MaxActiveConnectionIDs),
|
||||
}
|
||||
}
|
||||
|
||||
func (h *connIDManager) AddFromPreferredAddress(connID protocol.ConnectionID, resetToken protocol.StatelessResetToken) error {
|
||||
return h.addConnectionID(1, connID, resetToken)
|
||||
}
|
||||
|
||||
func (h *connIDManager) Add(f *wire.NewConnectionIDFrame) error {
|
||||
if err := h.add(f); err != nil {
|
||||
return err
|
||||
}
|
||||
if len(h.queue) >= protocol.MaxActiveConnectionIDs {
|
||||
return &qerr.TransportError{ErrorCode: qerr.ConnectionIDLimitError}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *connIDManager) add(f *wire.NewConnectionIDFrame) error {
|
||||
if h.activeConnectionID.Len() == 0 {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.ProtocolViolation,
|
||||
ErrorMessage: "received NEW_CONNECTION_ID frame but zero-length connection IDs are in use",
|
||||
}
|
||||
}
|
||||
// If the NEW_CONNECTION_ID frame is reordered, such that its sequence number is smaller than the currently active
|
||||
// connection ID or if it was already retired, send the RETIRE_CONNECTION_ID frame immediately.
|
||||
if f.SequenceNumber < max(h.activeSequenceNumber, h.highestProbingID) || f.SequenceNumber < h.highestRetired {
|
||||
h.queueControlFrame(&wire.RetireConnectionIDFrame{
|
||||
SequenceNumber: f.SequenceNumber,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
if f.RetirePriorTo != 0 && h.pathProbing != nil {
|
||||
for id, entry := range h.pathProbing {
|
||||
if entry.SequenceNumber < f.RetirePriorTo {
|
||||
h.queueControlFrame(&wire.RetireConnectionIDFrame{
|
||||
SequenceNumber: entry.SequenceNumber,
|
||||
})
|
||||
h.removeStatelessResetToken(entry.StatelessResetToken)
|
||||
delete(h.pathProbing, id)
|
||||
}
|
||||
}
|
||||
}
|
||||
// Retire elements in the queue.
|
||||
// Doesn't retire the active connection ID.
|
||||
if f.RetirePriorTo > h.highestRetired {
|
||||
var newQueue []newConnID
|
||||
for _, entry := range h.queue {
|
||||
if entry.SequenceNumber >= f.RetirePriorTo {
|
||||
newQueue = append(newQueue, entry)
|
||||
} else {
|
||||
h.queueControlFrame(&wire.RetireConnectionIDFrame{SequenceNumber: entry.SequenceNumber})
|
||||
}
|
||||
}
|
||||
h.queue = newQueue
|
||||
h.highestRetired = f.RetirePriorTo
|
||||
}
|
||||
|
||||
if f.SequenceNumber == h.activeSequenceNumber {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := h.addConnectionID(f.SequenceNumber, f.ConnectionID, f.StatelessResetToken); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Retire the active connection ID, if necessary.
|
||||
if h.activeSequenceNumber < f.RetirePriorTo {
|
||||
// The queue is guaranteed to have at least one element at this point.
|
||||
h.updateConnectionID()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *connIDManager) addConnectionID(seq uint64, connID protocol.ConnectionID, resetToken protocol.StatelessResetToken) error {
|
||||
// fast path: add to the end of the queue
|
||||
if len(h.queue) == 0 || h.queue[len(h.queue)-1].SequenceNumber < seq {
|
||||
h.queue = append(h.queue, newConnID{
|
||||
SequenceNumber: seq,
|
||||
ConnectionID: connID,
|
||||
StatelessResetToken: resetToken,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
// slow path: insert in the middle
|
||||
for i, entry := range h.queue {
|
||||
if entry.SequenceNumber == seq {
|
||||
if entry.ConnectionID != connID {
|
||||
return fmt.Errorf("received conflicting connection IDs for sequence number %d", seq)
|
||||
}
|
||||
if entry.StatelessResetToken != resetToken {
|
||||
return fmt.Errorf("received conflicting stateless reset tokens for sequence number %d", seq)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// insert at the correct position to maintain sorted order
|
||||
if entry.SequenceNumber > seq {
|
||||
h.queue = slices.Insert(h.queue, i, newConnID{
|
||||
SequenceNumber: seq,
|
||||
ConnectionID: connID,
|
||||
StatelessResetToken: resetToken,
|
||||
})
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return nil // unreachable
|
||||
}
|
||||
|
||||
func (h *connIDManager) updateConnectionID() {
|
||||
h.assertNotClosed()
|
||||
h.queueControlFrame(&wire.RetireConnectionIDFrame{
|
||||
SequenceNumber: h.activeSequenceNumber,
|
||||
})
|
||||
h.highestRetired = max(h.highestRetired, h.activeSequenceNumber)
|
||||
if h.activeStatelessResetToken != nil {
|
||||
h.removeStatelessResetToken(*h.activeStatelessResetToken)
|
||||
}
|
||||
|
||||
front := h.queue[0]
|
||||
h.queue = h.queue[1:]
|
||||
h.activeSequenceNumber = front.SequenceNumber
|
||||
h.activeConnectionID = front.ConnectionID
|
||||
h.activeStatelessResetToken = &front.StatelessResetToken
|
||||
h.packetsSinceLastChange = 0
|
||||
h.packetsPerConnectionID = protocol.PacketsPerConnectionID/2 + uint32(h.rand.Int31n(protocol.PacketsPerConnectionID))
|
||||
h.addStatelessResetToken(*h.activeStatelessResetToken)
|
||||
}
|
||||
|
||||
func (h *connIDManager) Close() {
|
||||
h.closed = true
|
||||
if h.activeStatelessResetToken != nil {
|
||||
h.removeStatelessResetToken(*h.activeStatelessResetToken)
|
||||
}
|
||||
if h.pathProbing != nil {
|
||||
for _, entry := range h.pathProbing {
|
||||
h.removeStatelessResetToken(entry.StatelessResetToken)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// is called when the server performs a Retry
|
||||
// and when the server changes the connection ID in the first Initial sent
|
||||
func (h *connIDManager) ChangeInitialConnID(newConnID protocol.ConnectionID) {
|
||||
if h.activeSequenceNumber != 0 {
|
||||
panic("expected first connection ID to have sequence number 0")
|
||||
}
|
||||
h.activeConnectionID = newConnID
|
||||
}
|
||||
|
||||
// is called when the server provides a stateless reset token in the transport parameters
|
||||
func (h *connIDManager) SetStatelessResetToken(token protocol.StatelessResetToken) {
|
||||
h.assertNotClosed()
|
||||
if h.activeSequenceNumber != 0 {
|
||||
panic("expected first connection ID to have sequence number 0")
|
||||
}
|
||||
h.activeStatelessResetToken = &token
|
||||
h.addStatelessResetToken(token)
|
||||
}
|
||||
|
||||
func (h *connIDManager) SentPacket() {
|
||||
h.packetsSinceLastChange++
|
||||
}
|
||||
|
||||
func (h *connIDManager) shouldUpdateConnID() bool {
|
||||
if !h.handshakeComplete {
|
||||
return false
|
||||
}
|
||||
// initiate the first change as early as possible (after handshake completion)
|
||||
if len(h.queue) > 0 && h.activeSequenceNumber == 0 {
|
||||
return true
|
||||
}
|
||||
// For later changes, only change if
|
||||
// 1. The queue of connection IDs is filled more than 50%.
|
||||
// 2. We sent at least PacketsPerConnectionID packets
|
||||
return 2*len(h.queue) >= protocol.MaxActiveConnectionIDs &&
|
||||
h.packetsSinceLastChange >= h.packetsPerConnectionID
|
||||
}
|
||||
|
||||
func (h *connIDManager) Get() protocol.ConnectionID {
|
||||
h.assertNotClosed()
|
||||
if h.shouldUpdateConnID() {
|
||||
h.updateConnectionID()
|
||||
}
|
||||
return h.activeConnectionID
|
||||
}
|
||||
|
||||
func (h *connIDManager) SetHandshakeComplete() {
|
||||
h.handshakeComplete = true
|
||||
}
|
||||
|
||||
// GetConnIDForPath retrieves a connection ID for a new path (i.e. not the active one).
|
||||
// Once a connection ID is allocated for a path, it cannot be used for a different path.
|
||||
// When called with the same pathID, it will return the same connection ID,
|
||||
// unless the peer requested that this connection ID be retired.
|
||||
func (h *connIDManager) GetConnIDForPath(id pathID) (protocol.ConnectionID, bool) {
|
||||
h.assertNotClosed()
|
||||
// if we're using zero-length connection IDs, we don't need to change the connection ID
|
||||
if h.activeConnectionID.Len() == 0 {
|
||||
return protocol.ConnectionID{}, true
|
||||
}
|
||||
|
||||
if h.pathProbing == nil {
|
||||
h.pathProbing = make(map[pathID]newConnID)
|
||||
}
|
||||
entry, ok := h.pathProbing[id]
|
||||
if ok {
|
||||
return entry.ConnectionID, true
|
||||
}
|
||||
if len(h.queue) == 0 {
|
||||
return protocol.ConnectionID{}, false
|
||||
}
|
||||
front := h.queue[0]
|
||||
h.queue = h.queue[1:]
|
||||
h.pathProbing[id] = front
|
||||
h.highestProbingID = front.SequenceNumber
|
||||
h.addStatelessResetToken(front.StatelessResetToken)
|
||||
return front.ConnectionID, true
|
||||
}
|
||||
|
||||
func (h *connIDManager) RetireConnIDForPath(pathID pathID) {
|
||||
h.assertNotClosed()
|
||||
// if we're using zero-length connection IDs, we don't need to change the connection ID
|
||||
if h.activeConnectionID.Len() == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
entry, ok := h.pathProbing[pathID]
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
h.queueControlFrame(&wire.RetireConnectionIDFrame{
|
||||
SequenceNumber: entry.SequenceNumber,
|
||||
})
|
||||
h.removeStatelessResetToken(entry.StatelessResetToken)
|
||||
delete(h.pathProbing, pathID)
|
||||
}
|
||||
|
||||
func (h *connIDManager) IsActiveStatelessResetToken(token protocol.StatelessResetToken) bool {
|
||||
if h.activeStatelessResetToken != nil {
|
||||
if *h.activeStatelessResetToken == token {
|
||||
return true
|
||||
}
|
||||
}
|
||||
if h.pathProbing != nil {
|
||||
for _, entry := range h.pathProbing {
|
||||
if entry.StatelessResetToken == token {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Using the connIDManager after it has been closed can have disastrous effects:
|
||||
// If the connection ID is rotated, a new entry would be inserted into the packet handler map,
|
||||
// leading to a memory leak of the connection struct.
|
||||
// See https://github.com/quic-go/quic-go/pull/4852 for more details.
|
||||
func (h *connIDManager) assertNotClosed() {
|
||||
if h.closed {
|
||||
panic("connection ID manager is closed")
|
||||
}
|
||||
}
|
||||
+3150
File diff suppressed because it is too large.
Load diff
+315
@@ -0,0 +1,315 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/internal/wire"
|
||||
"github.com/quic-go/quic-go/qlog"
|
||||
)
|
||||
|
||||
// ConvertFrame converts a wire.Frame into a logging.Frame.
|
||||
// This makes it possible for external packages to access the frames.
|
||||
// Furthermore, it removes the data slices from CRYPTO and STREAM frames.
|
||||
func toQlogFrame(frame wire.Frame) qlog.Frame {
|
||||
switch f := frame.(type) {
|
||||
case *wire.AckFrame:
|
||||
// We use a pool for ACK frames.
|
||||
// Implementations of the tracer interface may hold on to frames, so we need to make a copy here.
|
||||
return qlog.Frame{Frame: toQlogAckFrame(f)}
|
||||
case *wire.CryptoFrame:
|
||||
return qlog.Frame{
|
||||
Frame: &qlog.CryptoFrame{
|
||||
Offset: int64(f.Offset),
|
||||
Length: int64(len(f.Data)),
|
||||
},
|
||||
}
|
||||
case *wire.StreamFrame:
|
||||
return qlog.Frame{
|
||||
Frame: &qlog.StreamFrame{
|
||||
StreamID: f.StreamID,
|
||||
Offset: int64(f.Offset),
|
||||
Length: int64(f.DataLen()),
|
||||
Fin: f.Fin,
|
||||
},
|
||||
}
|
||||
case *wire.DatagramFrame:
|
||||
return qlog.Frame{
|
||||
Frame: &qlog.DatagramFrame{
|
||||
Length: int64(len(f.Data)),
|
||||
},
|
||||
}
|
||||
default:
|
||||
return qlog.Frame{Frame: frame}
|
||||
}
|
||||
}
|
||||
|
||||
func toQlogAckFrame(f *wire.AckFrame) *qlog.AckFrame {
|
||||
ack := &qlog.AckFrame{
|
||||
AckRanges: slices.Clone(f.AckRanges),
|
||||
DelayTime: f.DelayTime,
|
||||
ECNCE: f.ECNCE,
|
||||
ECT0: f.ECT0,
|
||||
ECT1: f.ECT1,
|
||||
}
|
||||
return ack
|
||||
}
|
||||
|
||||
func (c *Conn) logLongHeaderPacket(p *longHeaderPacket, ecn protocol.ECN, datagramPayloadChecksum qlog.DatagramPayloadChecksum) {
|
||||
// quic-go logging
|
||||
if c.logger.Debug() {
|
||||
p.header.Log(c.logger)
|
||||
if p.ack != nil {
|
||||
wire.LogFrame(c.logger, p.ack, true)
|
||||
}
|
||||
for _, frame := range p.frames {
|
||||
wire.LogFrame(c.logger, frame.Frame, true)
|
||||
}
|
||||
for _, frame := range p.streamFrames {
|
||||
wire.LogFrame(c.logger, frame.Frame, true)
|
||||
}
|
||||
}
|
||||
|
||||
// tracing
|
||||
if c.qlogger != nil {
|
||||
numFrames := len(p.frames) + len(p.streamFrames)
|
||||
if p.ack != nil {
|
||||
numFrames++
|
||||
}
|
||||
frames := make([]qlog.Frame, 0, numFrames)
|
||||
if p.ack != nil {
|
||||
frames = append(frames, toQlogFrame(p.ack))
|
||||
}
|
||||
for _, f := range p.frames {
|
||||
frames = append(frames, toQlogFrame(f.Frame))
|
||||
}
|
||||
for _, f := range p.streamFrames {
|
||||
frames = append(frames, toQlogFrame(f.Frame))
|
||||
}
|
||||
c.qlogger.RecordEvent(qlog.PacketSent{
|
||||
Header: qlog.PacketHeader{
|
||||
PacketType: toQlogPacketType(p.header.Type),
|
||||
KeyPhaseBit: p.header.KeyPhase,
|
||||
PacketNumber: p.header.PacketNumber,
|
||||
Version: p.header.Version,
|
||||
SrcConnectionID: p.header.SrcConnectionID,
|
||||
DestConnectionID: p.header.DestConnectionID,
|
||||
},
|
||||
Raw: qlog.RawInfo{
|
||||
Length: int(p.length),
|
||||
PayloadLength: int(p.header.Length),
|
||||
},
|
||||
DatagramPayloadChecksum: datagramPayloadChecksum,
|
||||
Frames: frames,
|
||||
ECN: toQlogECN(ecn),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) logShortHeaderPacket(p shortHeaderPacket, ecn protocol.ECN, size protocol.ByteCount) {
|
||||
c.logShortHeaderPacketWithDatagramPayloadChecksum(p, ecn, size, false, 0)
|
||||
}
|
||||
|
||||
func (c *Conn) logShortHeaderPacketWithDatagramPayloadChecksum(p shortHeaderPacket, ecn protocol.ECN, size protocol.ByteCount, isCoalesced bool, datagramPayloadChecksum qlog.DatagramPayloadChecksum) {
|
||||
if c.logger.Debug() && !isCoalesced {
|
||||
c.logger.Debugf("-> Sending packet %d (%d bytes) for connection %s, 1-RTT (ECN: %s)", p.PacketNumber, size, c.logID, ecn)
|
||||
}
|
||||
// quic-go logging
|
||||
if c.logger.Debug() {
|
||||
wire.LogShortHeader(c.logger, p.DestConnID, p.PacketNumber, p.PacketNumberLen, p.KeyPhase)
|
||||
if p.Ack != nil {
|
||||
wire.LogFrame(c.logger, p.Ack, true)
|
||||
}
|
||||
for _, f := range p.Frames {
|
||||
wire.LogFrame(c.logger, f.Frame, true)
|
||||
}
|
||||
for _, f := range p.StreamFrames {
|
||||
wire.LogFrame(c.logger, f.Frame, true)
|
||||
}
|
||||
}
|
||||
|
||||
// tracing
|
||||
if c.qlogger != nil {
|
||||
numFrames := len(p.Frames) + len(p.StreamFrames)
|
||||
if p.Ack != nil {
|
||||
numFrames++
|
||||
}
|
||||
fs := make([]qlog.Frame, 0, numFrames)
|
||||
if p.Ack != nil {
|
||||
fs = append(fs, toQlogFrame(p.Ack))
|
||||
}
|
||||
for _, f := range p.Frames {
|
||||
fs = append(fs, toQlogFrame(f.Frame))
|
||||
}
|
||||
for _, f := range p.StreamFrames {
|
||||
fs = append(fs, toQlogFrame(f.Frame))
|
||||
}
|
||||
c.qlogger.RecordEvent(qlog.PacketSent{
|
||||
Header: qlog.PacketHeader{
|
||||
PacketType: qlog.PacketType1RTT,
|
||||
KeyPhaseBit: p.KeyPhase,
|
||||
PacketNumber: p.PacketNumber,
|
||||
Version: c.version,
|
||||
DestConnectionID: p.DestConnID,
|
||||
},
|
||||
Raw: qlog.RawInfo{
|
||||
Length: int(size),
|
||||
PayloadLength: int(size - wire.ShortHeaderLen(p.DestConnID, p.PacketNumberLen)),
|
||||
},
|
||||
DatagramPayloadChecksum: datagramPayloadChecksum,
|
||||
Frames: fs,
|
||||
ECN: toQlogECN(ecn),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) logCoalescedPacket(packet *coalescedPacket, ecn protocol.ECN) {
|
||||
var datagramPayloadChecksum qlog.DatagramPayloadChecksum
|
||||
if c.qlogger != nil {
|
||||
datagramPayloadChecksum = qlog.CalculateDatagramPayloadChecksum(packet.buffer.Data)
|
||||
}
|
||||
if c.logger.Debug() {
|
||||
// There's a short period between dropping both Initial and Handshake keys and completion of the handshake,
|
||||
// during which we might call PackCoalescedPacket but just pack a short header packet.
|
||||
if len(packet.longHdrPackets) == 0 && packet.shortHdrPacket != nil {
|
||||
c.logShortHeaderPacketWithDatagramPayloadChecksum(
|
||||
*packet.shortHdrPacket,
|
||||
ecn,
|
||||
packet.shortHdrPacket.Length,
|
||||
false,
|
||||
datagramPayloadChecksum,
|
||||
)
|
||||
return
|
||||
}
|
||||
if len(packet.longHdrPackets) > 1 {
|
||||
c.logger.Debugf("-> Sending coalesced packet (%d parts, %d bytes) for connection %s", len(packet.longHdrPackets), packet.buffer.Len(), c.logID)
|
||||
} else {
|
||||
c.logger.Debugf("-> Sending packet %d (%d bytes) for connection %s, %s", packet.longHdrPackets[0].header.PacketNumber, packet.buffer.Len(), c.logID, packet.longHdrPackets[0].EncryptionLevel())
|
||||
}
|
||||
}
|
||||
for _, p := range packet.longHdrPackets {
|
||||
c.logLongHeaderPacket(p, ecn, datagramPayloadChecksum)
|
||||
}
|
||||
if p := packet.shortHdrPacket; p != nil {
|
||||
c.logShortHeaderPacketWithDatagramPayloadChecksum(*p, ecn, p.Length, true, datagramPayloadChecksum)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) qlogTransportParameters(tp *wire.TransportParameters, sentBy protocol.Perspective, restore bool) {
|
||||
ev := qlog.ParametersSet{
|
||||
Restore: restore,
|
||||
OriginalDestinationConnectionID: tp.OriginalDestinationConnectionID,
|
||||
InitialSourceConnectionID: tp.InitialSourceConnectionID,
|
||||
RetrySourceConnectionID: tp.RetrySourceConnectionID,
|
||||
StatelessResetToken: tp.StatelessResetToken,
|
||||
DisableActiveMigration: tp.DisableActiveMigration,
|
||||
MaxIdleTimeout: tp.MaxIdleTimeout,
|
||||
MaxUDPPayloadSize: tp.MaxUDPPayloadSize,
|
||||
AckDelayExponent: tp.AckDelayExponent,
|
||||
MaxAckDelay: tp.MaxAckDelay,
|
||||
ActiveConnectionIDLimit: tp.ActiveConnectionIDLimit,
|
||||
InitialMaxData: tp.InitialMaxData,
|
||||
InitialMaxStreamDataBidiLocal: tp.InitialMaxStreamDataBidiLocal,
|
||||
InitialMaxStreamDataBidiRemote: tp.InitialMaxStreamDataBidiRemote,
|
||||
InitialMaxStreamDataUni: tp.InitialMaxStreamDataUni,
|
||||
InitialMaxStreamsBidi: int64(tp.MaxBidiStreamNum),
|
||||
InitialMaxStreamsUni: int64(tp.MaxUniStreamNum),
|
||||
MaxDatagramFrameSize: tp.MaxDatagramFrameSize,
|
||||
EnableResetStreamAt: tp.EnableResetStreamAt,
|
||||
}
|
||||
if sentBy == c.perspective {
|
||||
ev.Initiator = qlog.InitiatorLocal
|
||||
} else {
|
||||
ev.Initiator = qlog.InitiatorRemote
|
||||
}
|
||||
if tp.PreferredAddress != nil {
|
||||
ev.PreferredAddress = &qlog.PreferredAddress{
|
||||
IPv4: tp.PreferredAddress.IPv4,
|
||||
IPv6: tp.PreferredAddress.IPv6,
|
||||
ConnectionID: tp.PreferredAddress.ConnectionID,
|
||||
StatelessResetToken: tp.PreferredAddress.StatelessResetToken,
|
||||
}
|
||||
}
|
||||
c.qlogger.RecordEvent(ev)
|
||||
}
|
||||
|
||||
func toQlogECN(ecn protocol.ECN) qlog.ECN {
|
||||
//nolint:exhaustive // only need to handle the 3 valid values
|
||||
switch ecn {
|
||||
case protocol.ECT0:
|
||||
return qlog.ECT0
|
||||
case protocol.ECT1:
|
||||
return qlog.ECT1
|
||||
case protocol.ECNCE:
|
||||
return qlog.ECNCE
|
||||
default:
|
||||
return qlog.ECNUnsupported
|
||||
}
|
||||
}
|
||||
|
||||
func toQlogPacketType(pt protocol.PacketType) qlog.PacketType {
|
||||
var qpt qlog.PacketType
|
||||
switch pt {
|
||||
case protocol.PacketTypeInitial:
|
||||
qpt = qlog.PacketTypeInitial
|
||||
case protocol.PacketTypeHandshake:
|
||||
qpt = qlog.PacketTypeHandshake
|
||||
case protocol.PacketType0RTT:
|
||||
qpt = qlog.PacketType0RTT
|
||||
case protocol.PacketTypeRetry:
|
||||
qpt = qlog.PacketTypeRetry
|
||||
}
|
||||
return qpt
|
||||
}
|
||||
|
||||
func toPathEndpointInfo(addr *net.UDPAddr) qlog.PathEndpointInfo {
|
||||
if addr == nil {
|
||||
return qlog.PathEndpointInfo{}
|
||||
}
|
||||
|
||||
var info qlog.PathEndpointInfo
|
||||
if addr.IP == nil || addr.IP.To4() != nil {
|
||||
addrPort := netip.AddrPortFrom(netip.AddrFrom4([4]byte(addr.IP.To4())), uint16(addr.Port))
|
||||
if addrPort.IsValid() {
|
||||
info.IPv4 = addrPort
|
||||
}
|
||||
} else {
|
||||
addrPort := netip.AddrPortFrom(netip.AddrFrom16([16]byte(addr.IP.To16())), uint16(addr.Port))
|
||||
if addrPort.IsValid() {
|
||||
info.IPv6 = addrPort
|
||||
}
|
||||
}
|
||||
return info
|
||||
}
|
||||
|
||||
// startedConnectionEvent builds a StartedConnection event using consistent logic
|
||||
// for both endpoints. If the local address is unspecified (e.g., dual-stack
|
||||
// listener), it selects the family based on the remote address and uses the
|
||||
// unspecified address of that family with the local port.
|
||||
func startedConnectionEvent(local, remote *net.UDPAddr) qlog.StartedConnection {
|
||||
var localInfo, remoteInfo qlog.PathEndpointInfo
|
||||
if remote != nil {
|
||||
remoteInfo = toPathEndpointInfo(remote)
|
||||
}
|
||||
if local != nil {
|
||||
if local.IP == nil || local.IP.IsUnspecified() {
|
||||
// Choose local family based on the remote address family.
|
||||
if remote != nil && remote.IP.To4() != nil {
|
||||
ap := netip.AddrPortFrom(netip.AddrFrom4([4]byte{}), uint16(local.Port))
|
||||
if ap.IsValid() {
|
||||
localInfo.IPv4 = ap
|
||||
}
|
||||
} else if remote != nil && remote.IP.To16() != nil && remote.IP.To4() == nil {
|
||||
ap := netip.AddrPortFrom(netip.AddrFrom16([16]byte{}), uint16(local.Port))
|
||||
if ap.IsValid() {
|
||||
localInfo.IPv6 = ap
|
||||
}
|
||||
}
|
||||
} else {
|
||||
localInfo = toPathEndpointInfo(local)
|
||||
}
|
||||
}
|
||||
return qlog.StartedConnection{Local: localInfo, Remote: remoteInfo}
|
||||
}
|
||||
+249
@@ -0,0 +1,249 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"slices"
|
||||
"strconv"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/internal/qerr"
|
||||
"github.com/quic-go/quic-go/internal/wire"
|
||||
)
|
||||
|
||||
const disableClientHelloScramblingEnv = "QUIC_GO_DISABLE_CLIENTHELLO_SCRAMBLING"
|
||||
|
||||
// The baseCryptoStream is used by the cryptoStream and the initialCryptoStream.
|
||||
// This allows us to implement different logic for PopCryptoFrame for the two streams.
|
||||
type baseCryptoStream struct {
|
||||
queue frameSorter
|
||||
|
||||
highestOffset protocol.ByteCount
|
||||
finished bool
|
||||
|
||||
writeOffset protocol.ByteCount
|
||||
writeBuf []byte
|
||||
}
|
||||
|
||||
func newCryptoStream() *cryptoStream {
|
||||
return &cryptoStream{baseCryptoStream{queue: *newFrameSorter()}}
|
||||
}
|
||||
|
||||
func (s *baseCryptoStream) HandleCryptoFrame(f *wire.CryptoFrame) error {
|
||||
highestOffset := f.Offset + protocol.ByteCount(len(f.Data))
|
||||
if maxOffset := highestOffset; maxOffset > protocol.MaxCryptoStreamOffset {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.CryptoBufferExceeded,
|
||||
ErrorMessage: fmt.Sprintf("received invalid offset %d on crypto stream, maximum allowed %d", maxOffset, protocol.MaxCryptoStreamOffset),
|
||||
}
|
||||
}
|
||||
if s.finished {
|
||||
if highestOffset > s.highestOffset {
|
||||
// reject crypto data received after this stream was already finished
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.ProtocolViolation,
|
||||
ErrorMessage: "received crypto data after change of encryption level",
|
||||
}
|
||||
}
|
||||
// ignore data with a smaller offset than the highest received
|
||||
// could e.g. be a retransmission
|
||||
return nil
|
||||
}
|
||||
s.highestOffset = max(s.highestOffset, highestOffset)
|
||||
return s.queue.Push(f.Data, f.Offset, nil)
|
||||
}
|
||||
|
||||
// GetCryptoData retrieves data that was received in CRYPTO frames
|
||||
func (s *baseCryptoStream) GetCryptoData() []byte {
|
||||
_, data, _ := s.queue.Pop()
|
||||
return data
|
||||
}
|
||||
|
||||
func (s *baseCryptoStream) Finish() error {
|
||||
if s.queue.HasMoreData() {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.ProtocolViolation,
|
||||
ErrorMessage: "encryption level changed, but crypto stream has more data to read",
|
||||
}
|
||||
}
|
||||
s.finished = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// Writes writes data that should be sent out in CRYPTO frames
|
||||
func (s *baseCryptoStream) Write(p []byte) (int, error) {
|
||||
s.writeBuf = append(s.writeBuf, p...)
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func (s *baseCryptoStream) HasData() bool {
|
||||
return len(s.writeBuf) > 0
|
||||
}
|
||||
|
||||
func (s *baseCryptoStream) PopCryptoFrame(maxLen protocol.ByteCount) *wire.CryptoFrame {
|
||||
f := &wire.CryptoFrame{Offset: s.writeOffset}
|
||||
n := min(f.MaxDataLen(maxLen), protocol.ByteCount(len(s.writeBuf)))
|
||||
if n <= 0 {
|
||||
return nil
|
||||
}
|
||||
f.Data = s.writeBuf[:n]
|
||||
s.writeBuf = s.writeBuf[n:]
|
||||
s.writeOffset += n
|
||||
return f
|
||||
}
|
||||
|
||||
type cryptoStream struct {
|
||||
baseCryptoStream
|
||||
}
|
||||
|
||||
type clientHelloCut struct {
|
||||
start protocol.ByteCount
|
||||
end protocol.ByteCount
|
||||
}
|
||||
|
||||
type initialCryptoStream struct {
|
||||
baseCryptoStream
|
||||
|
||||
scramble bool
|
||||
end protocol.ByteCount
|
||||
cuts [2]clientHelloCut
|
||||
}
|
||||
|
||||
func newInitialCryptoStream(isClient bool) *initialCryptoStream {
|
||||
var scramble bool
|
||||
if isClient {
|
||||
disabled, err := strconv.ParseBool(os.Getenv(disableClientHelloScramblingEnv))
|
||||
scramble = err != nil || !disabled
|
||||
}
|
||||
s := &initialCryptoStream{
|
||||
baseCryptoStream: baseCryptoStream{queue: *newFrameSorter()},
|
||||
scramble: scramble,
|
||||
}
|
||||
for i := range len(s.cuts) {
|
||||
s.cuts[i].start = protocol.InvalidByteCount
|
||||
s.cuts[i].end = protocol.InvalidByteCount
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func (s *initialCryptoStream) HasData() bool {
|
||||
// The ClientHello might be written in multiple parts.
|
||||
// In order to correctly split the ClientHello, we need the entire ClientHello has been queued.
|
||||
if s.scramble && s.writeOffset == 0 && s.cuts[0].start == protocol.InvalidByteCount {
|
||||
return false
|
||||
}
|
||||
return s.baseCryptoStream.HasData()
|
||||
}
|
||||
|
||||
func (s *initialCryptoStream) Write(p []byte) (int, error) {
|
||||
s.writeBuf = append(s.writeBuf, p...)
|
||||
if !s.scramble {
|
||||
return len(p), nil
|
||||
}
|
||||
if s.cuts[0].start == protocol.InvalidByteCount {
|
||||
sniPos, sniLen, echPos, err := findSNIAndECH(s.writeBuf)
|
||||
if errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
return len(p), nil
|
||||
}
|
||||
if err != nil {
|
||||
return len(p), err
|
||||
}
|
||||
if sniPos == -1 && echPos == -1 {
|
||||
// Neither SNI nor ECH found.
|
||||
// There's nothing to scramble.
|
||||
s.scramble = false
|
||||
return len(p), nil
|
||||
}
|
||||
s.end = protocol.ByteCount(len(s.writeBuf))
|
||||
s.cuts[0].start = protocol.ByteCount(sniPos + sniLen/2) // right in the middle
|
||||
s.cuts[0].end = protocol.ByteCount(sniPos + sniLen)
|
||||
if echPos > 0 {
|
||||
// ECH extension found, cut the ECH extension type value (a uint16) in half
|
||||
start := protocol.ByteCount(echPos + 1)
|
||||
s.cuts[1].start = start
|
||||
// cut somewhere (16 bytes), most likely in the ECH extension value
|
||||
s.cuts[1].end = min(start+16, s.end)
|
||||
}
|
||||
slices.SortFunc(s.cuts[:], func(a, b clientHelloCut) int {
|
||||
if a.start == protocol.InvalidByteCount {
|
||||
return 1
|
||||
}
|
||||
if a.start > b.start {
|
||||
return 1
|
||||
}
|
||||
return -1
|
||||
})
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func (s *initialCryptoStream) PopCryptoFrame(maxLen protocol.ByteCount) *wire.CryptoFrame {
|
||||
if !s.scramble {
|
||||
return s.baseCryptoStream.PopCryptoFrame(maxLen)
|
||||
}
|
||||
|
||||
// send out the skipped parts
|
||||
if s.writeOffset == s.end {
|
||||
var foundCuts bool
|
||||
var f *wire.CryptoFrame
|
||||
for i, c := range s.cuts {
|
||||
if c.start == protocol.InvalidByteCount {
|
||||
continue
|
||||
}
|
||||
foundCuts = true
|
||||
if f != nil {
|
||||
break
|
||||
}
|
||||
f = &wire.CryptoFrame{Offset: c.start}
|
||||
n := min(f.MaxDataLen(maxLen), c.end-c.start)
|
||||
if n <= 0 {
|
||||
return nil
|
||||
}
|
||||
f.Data = s.writeBuf[c.start : c.start+n]
|
||||
s.cuts[i].start += n
|
||||
if s.cuts[i].start == c.end {
|
||||
s.cuts[i].start = protocol.InvalidByteCount
|
||||
s.cuts[i].end = protocol.InvalidByteCount
|
||||
foundCuts = false
|
||||
}
|
||||
}
|
||||
if !foundCuts {
|
||||
// no more cuts found, we're done sending out everything up until s.end
|
||||
s.writeBuf = s.writeBuf[s.end:]
|
||||
s.end = protocol.InvalidByteCount
|
||||
s.scramble = false
|
||||
}
|
||||
return f
|
||||
}
|
||||
|
||||
nextCut := clientHelloCut{start: protocol.InvalidByteCount, end: protocol.InvalidByteCount}
|
||||
for _, c := range s.cuts {
|
||||
if c.start == protocol.InvalidByteCount {
|
||||
continue
|
||||
}
|
||||
if c.start > s.writeOffset {
|
||||
nextCut = c
|
||||
break
|
||||
}
|
||||
}
|
||||
f := &wire.CryptoFrame{Offset: s.writeOffset}
|
||||
maxOffset := nextCut.start
|
||||
if maxOffset == protocol.InvalidByteCount {
|
||||
maxOffset = s.end
|
||||
}
|
||||
n := min(f.MaxDataLen(maxLen), maxOffset-s.writeOffset)
|
||||
if n <= 0 {
|
||||
return nil
|
||||
}
|
||||
f.Data = s.writeBuf[s.writeOffset : s.writeOffset+n]
|
||||
// Don't reslice the writeBuf yet.
|
||||
// This is done once all parts have been sent out.
|
||||
s.writeOffset += n
|
||||
if s.writeOffset == nextCut.start {
|
||||
s.writeOffset = nextCut.end
|
||||
}
|
||||
|
||||
return f
|
||||
}
|
||||
+73
@@ -0,0 +1,73 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/internal/wire"
|
||||
)
|
||||
|
||||
type cryptoStreamManager struct {
|
||||
initialStream *initialCryptoStream
|
||||
handshakeStream *cryptoStream
|
||||
oneRTTStream *cryptoStream
|
||||
}
|
||||
|
||||
func newCryptoStreamManager(
|
||||
initialStream *initialCryptoStream,
|
||||
handshakeStream *cryptoStream,
|
||||
oneRTTStream *cryptoStream,
|
||||
) *cryptoStreamManager {
|
||||
return &cryptoStreamManager{
|
||||
initialStream: initialStream,
|
||||
handshakeStream: handshakeStream,
|
||||
oneRTTStream: oneRTTStream,
|
||||
}
|
||||
}
|
||||
|
||||
func (m *cryptoStreamManager) HandleCryptoFrame(frame *wire.CryptoFrame, encLevel protocol.EncryptionLevel) error {
|
||||
//nolint:exhaustive // CRYPTO frames cannot be sent in 0-RTT packets.
|
||||
switch encLevel {
|
||||
case protocol.EncryptionInitial:
|
||||
return m.initialStream.HandleCryptoFrame(frame)
|
||||
case protocol.EncryptionHandshake:
|
||||
return m.handshakeStream.HandleCryptoFrame(frame)
|
||||
case protocol.Encryption1RTT:
|
||||
return m.oneRTTStream.HandleCryptoFrame(frame)
|
||||
default:
|
||||
return fmt.Errorf("received CRYPTO frame with unexpected encryption level: %s", encLevel)
|
||||
}
|
||||
}
|
||||
|
||||
func (m *cryptoStreamManager) GetCryptoData(encLevel protocol.EncryptionLevel) []byte {
|
||||
//nolint:exhaustive // CRYPTO frames cannot be sent in 0-RTT packets.
|
||||
switch encLevel {
|
||||
case protocol.EncryptionInitial:
|
||||
return m.initialStream.GetCryptoData()
|
||||
case protocol.EncryptionHandshake:
|
||||
return m.handshakeStream.GetCryptoData()
|
||||
case protocol.Encryption1RTT:
|
||||
return m.oneRTTStream.GetCryptoData()
|
||||
default:
|
||||
panic(fmt.Sprintf("received CRYPTO frame with unexpected encryption level: %s", encLevel))
|
||||
}
|
||||
}
|
||||
|
||||
func (m *cryptoStreamManager) GetPostHandshakeData(maxSize protocol.ByteCount) *wire.CryptoFrame {
|
||||
if !m.oneRTTStream.HasData() {
|
||||
return nil
|
||||
}
|
||||
return m.oneRTTStream.PopCryptoFrame(maxSize)
|
||||
}
|
||||
|
||||
func (m *cryptoStreamManager) Drop(encLevel protocol.EncryptionLevel) error {
|
||||
//nolint:exhaustive // 1-RTT keys should never get dropped.
|
||||
switch encLevel {
|
||||
case protocol.EncryptionInitial:
|
||||
return m.initialStream.Finish()
|
||||
case protocol.EncryptionHandshake:
|
||||
return m.handshakeStream.Finish()
|
||||
default:
|
||||
panic(fmt.Sprintf("dropped unexpected encryption level: %s", encLevel))
|
||||
}
|
||||
}
|
||||
+137
@@ -0,0 +1,137 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/utils"
|
||||
"github.com/quic-go/quic-go/internal/utils/ringbuffer"
|
||||
"github.com/quic-go/quic-go/internal/wire"
|
||||
)
|
||||
|
||||
const (
|
||||
maxDatagramSendQueueLen = 32
|
||||
maxDatagramRcvQueueLen = 128
|
||||
)
|
||||
|
||||
type datagramQueue struct {
|
||||
sendMx sync.Mutex
|
||||
sendQueue ringbuffer.RingBuffer[*wire.DatagramFrame]
|
||||
sent chan struct{} // used to notify Add that a datagram was dequeued
|
||||
|
||||
rcvMx sync.Mutex
|
||||
rcvQueue [][]byte
|
||||
rcvd chan struct{} // used to notify Receive that a new datagram was received
|
||||
|
||||
closeErr error
|
||||
closed chan struct{}
|
||||
|
||||
hasData func()
|
||||
|
||||
logger utils.Logger
|
||||
}
|
||||
|
||||
func newDatagramQueue(hasData func(), logger utils.Logger) *datagramQueue {
|
||||
return &datagramQueue{
|
||||
hasData: hasData,
|
||||
rcvd: make(chan struct{}, 1),
|
||||
sent: make(chan struct{}, 1),
|
||||
closed: make(chan struct{}),
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
// Add queues a new DATAGRAM frame for sending.
|
||||
// Up to 32 DATAGRAM frames will be queued.
|
||||
// Once that limit is reached, Add blocks until the queue size has reduced.
|
||||
func (h *datagramQueue) Add(f *wire.DatagramFrame) error {
|
||||
h.sendMx.Lock()
|
||||
|
||||
for {
|
||||
if h.sendQueue.Len() < maxDatagramSendQueueLen {
|
||||
h.sendQueue.PushBack(f)
|
||||
h.sendMx.Unlock()
|
||||
h.hasData()
|
||||
return nil
|
||||
}
|
||||
select {
|
||||
case <-h.sent: // drain the queue so we don't loop immediately
|
||||
default:
|
||||
}
|
||||
h.sendMx.Unlock()
|
||||
select {
|
||||
case <-h.closed:
|
||||
return h.closeErr
|
||||
case <-h.sent:
|
||||
}
|
||||
h.sendMx.Lock()
|
||||
}
|
||||
}
|
||||
|
||||
// Peek gets the next DATAGRAM frame for sending.
|
||||
// If actually sent out, Pop needs to be called before the next call to Peek.
|
||||
func (h *datagramQueue) Peek() *wire.DatagramFrame {
|
||||
h.sendMx.Lock()
|
||||
defer h.sendMx.Unlock()
|
||||
if h.sendQueue.Empty() {
|
||||
return nil
|
||||
}
|
||||
return h.sendQueue.PeekFront()
|
||||
}
|
||||
|
||||
func (h *datagramQueue) Pop() {
|
||||
h.sendMx.Lock()
|
||||
defer h.sendMx.Unlock()
|
||||
_ = h.sendQueue.PopFront()
|
||||
select {
|
||||
case h.sent <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
// HandleDatagramFrame handles a received DATAGRAM frame.
|
||||
func (h *datagramQueue) HandleDatagramFrame(f *wire.DatagramFrame) {
|
||||
data := make([]byte, len(f.Data))
|
||||
copy(data, f.Data)
|
||||
var queued bool
|
||||
h.rcvMx.Lock()
|
||||
if len(h.rcvQueue) < maxDatagramRcvQueueLen {
|
||||
h.rcvQueue = append(h.rcvQueue, data)
|
||||
queued = true
|
||||
select {
|
||||
case h.rcvd <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
h.rcvMx.Unlock()
|
||||
if !queued && h.logger.Debug() {
|
||||
h.logger.Debugf("Discarding received DATAGRAM frame (%d bytes payload)", len(f.Data))
|
||||
}
|
||||
}
|
||||
|
||||
// Receive gets a received DATAGRAM frame.
|
||||
func (h *datagramQueue) Receive(ctx context.Context) ([]byte, error) {
|
||||
for {
|
||||
h.rcvMx.Lock()
|
||||
if len(h.rcvQueue) > 0 {
|
||||
data := h.rcvQueue[0]
|
||||
h.rcvQueue = h.rcvQueue[1:]
|
||||
h.rcvMx.Unlock()
|
||||
return data, nil
|
||||
}
|
||||
h.rcvMx.Unlock()
|
||||
select {
|
||||
case <-h.rcvd:
|
||||
continue
|
||||
case <-h.closed:
|
||||
return nil, h.closeErr
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *datagramQueue) CloseWithError(e error) {
|
||||
h.closeErr = e
|
||||
close(h.closed)
|
||||
}
|
||||
+105
@@ -0,0 +1,105 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/qerr"
|
||||
)
|
||||
|
||||
type (
|
||||
// TransportError indicates an error that occurred on the QUIC transport layer.
|
||||
// Every transport error other than CONNECTION_REFUSED and APPLICATION_ERROR is
|
||||
// likely a bug in the implementation.
|
||||
TransportError = qerr.TransportError
|
||||
// ApplicationError is an application-defined error.
|
||||
ApplicationError = qerr.ApplicationError
|
||||
// VersionNegotiationError indicates a failure to negotiate a QUIC version.
|
||||
VersionNegotiationError = qerr.VersionNegotiationError
|
||||
// StatelessResetError indicates a stateless reset was received.
|
||||
// This can happen when the peer reboots, or when packets are misrouted.
|
||||
// See section 10.3 of RFC 9000 for details.
|
||||
StatelessResetError = qerr.StatelessResetError
|
||||
// IdleTimeoutError indicates that the connection timed out because it was inactive for too long.
|
||||
IdleTimeoutError = qerr.IdleTimeoutError
|
||||
// HandshakeTimeoutError indicates that the connection timed out before completing the handshake.
|
||||
HandshakeTimeoutError = qerr.HandshakeTimeoutError
|
||||
)
|
||||
|
||||
type (
|
||||
// TransportErrorCode is a QUIC transport error code, see section 20 of RFC 9000.
|
||||
TransportErrorCode = qerr.TransportErrorCode
|
||||
// ApplicationErrorCode is an QUIC application error code.
|
||||
ApplicationErrorCode = qerr.ApplicationErrorCode
|
||||
// StreamErrorCode is a QUIC stream error code. The meaning of the value is defined by the application.
|
||||
StreamErrorCode = qerr.StreamErrorCode
|
||||
)
|
||||
|
||||
const (
|
||||
// NoError is the NO_ERROR transport error code.
|
||||
NoError = qerr.NoError
|
||||
// InternalError is the INTERNAL_ERROR transport error code.
|
||||
InternalError = qerr.InternalError
|
||||
// ConnectionRefused is the CONNECTION_REFUSED transport error code.
|
||||
ConnectionRefused = qerr.ConnectionRefused
|
||||
// FlowControlError is the FLOW_CONTROL_ERROR transport error code.
|
||||
FlowControlError = qerr.FlowControlError
|
||||
// StreamLimitError is the STREAM_LIMIT_ERROR transport error code.
|
||||
StreamLimitError = qerr.StreamLimitError
|
||||
// StreamStateError is the STREAM_STATE_ERROR transport error code.
|
||||
StreamStateError = qerr.StreamStateError
|
||||
// FinalSizeError is the FINAL_SIZE_ERROR transport error code.
|
||||
FinalSizeError = qerr.FinalSizeError
|
||||
// FrameEncodingError is the FRAME_ENCODING_ERROR transport error code.
|
||||
FrameEncodingError = qerr.FrameEncodingError
|
||||
// TransportParameterError is the TRANSPORT_PARAMETER_ERROR transport error code.
|
||||
TransportParameterError = qerr.TransportParameterError
|
||||
// ConnectionIDLimitError is the CONNECTION_ID_LIMIT_ERROR transport error code.
|
||||
ConnectionIDLimitError = qerr.ConnectionIDLimitError
|
||||
// ProtocolViolation is the PROTOCOL_VIOLATION transport error code.
|
||||
ProtocolViolation = qerr.ProtocolViolation
|
||||
// InvalidToken is the INVALID_TOKEN transport error code.
|
||||
InvalidToken = qerr.InvalidToken
|
||||
// ApplicationErrorErrorCode is the APPLICATION_ERROR transport error code.
|
||||
ApplicationErrorErrorCode = qerr.ApplicationErrorErrorCode
|
||||
// CryptoBufferExceeded is the CRYPTO_BUFFER_EXCEEDED transport error code.
|
||||
CryptoBufferExceeded = qerr.CryptoBufferExceeded
|
||||
// KeyUpdateError is the KEY_UPDATE_ERROR transport error code.
|
||||
KeyUpdateError = qerr.KeyUpdateError
|
||||
// AEADLimitReached is the AEAD_LIMIT_REACHED transport error code.
|
||||
AEADLimitReached = qerr.AEADLimitReached
|
||||
// NoViablePathError is the NO_VIABLE_PATH_ERROR transport error code.
|
||||
NoViablePathError = qerr.NoViablePathError
|
||||
)
|
||||
|
||||
// A StreamError is used to signal stream cancellations.
|
||||
// It is returned from the Read and Write methods of the [ReceiveStream], [SendStream] and [Stream].
|
||||
type StreamError struct {
|
||||
StreamID StreamID
|
||||
ErrorCode StreamErrorCode
|
||||
Remote bool
|
||||
}
|
||||
|
||||
func (e *StreamError) Is(target error) bool {
|
||||
t, ok := target.(*StreamError)
|
||||
return ok && e.StreamID == t.StreamID && e.ErrorCode == t.ErrorCode && e.Remote == t.Remote
|
||||
}
|
||||
|
||||
func (e *StreamError) Error() string {
|
||||
pers := "local"
|
||||
if e.Remote {
|
||||
pers = "remote"
|
||||
}
|
||||
return fmt.Sprintf("stream %d canceled by %s with error code %d", e.StreamID, pers, e.ErrorCode)
|
||||
}
|
||||
|
||||
// DatagramTooLargeError is returned from Conn.SendDatagram if the payload is too large to be sent.
|
||||
type DatagramTooLargeError struct {
|
||||
MaxDatagramPayloadSize int64
|
||||
}
|
||||
|
||||
func (e *DatagramTooLargeError) Is(target error) bool {
|
||||
t, ok := target.(*DatagramTooLargeError)
|
||||
return ok && e.MaxDatagramPayloadSize == t.MaxDatagramPayloadSize
|
||||
}
|
||||
|
||||
func (e *DatagramTooLargeError) Error() string { return "DATAGRAM frame too large" }
|
||||
+84
@@ -0,0 +1,84 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/monotime"
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/internal/utils"
|
||||
)
|
||||
|
||||
type receiveFlowController struct {
|
||||
//nolint:structcheck // The mutex is used both by the stream and the connection flow controller
|
||||
mutex sync.Mutex
|
||||
bytesRead protocol.ByteCount
|
||||
highestReceived protocol.ByteCount
|
||||
receiveWindow protocol.ByteCount
|
||||
receiveWindowSize protocol.ByteCount
|
||||
maxReceiveWindowSize protocol.ByteCount
|
||||
|
||||
allowWindowIncrease func(size protocol.ByteCount) bool
|
||||
|
||||
epochStartTime monotime.Time
|
||||
epochStartOffset protocol.ByteCount
|
||||
rttStats *utils.RTTStats
|
||||
|
||||
logger utils.Logger
|
||||
}
|
||||
|
||||
// needs to be called with locked mutex
|
||||
func (c *receiveFlowController) addBytesRead(n protocol.ByteCount) {
|
||||
c.bytesRead += n
|
||||
}
|
||||
|
||||
func (c *receiveFlowController) hasWindowUpdate() bool {
|
||||
bytesRemaining := c.receiveWindow - c.bytesRead
|
||||
// update the window when more than the threshold was consumed
|
||||
return bytesRemaining <= protocol.ByteCount(float64(c.receiveWindowSize)*(1-protocol.WindowUpdateThreshold))
|
||||
}
|
||||
|
||||
// getWindowUpdate updates the receive window, if necessary
|
||||
// it returns the new offset
|
||||
func (c *receiveFlowController) getWindowUpdate(now monotime.Time) protocol.ByteCount {
|
||||
if !c.hasWindowUpdate() {
|
||||
return 0
|
||||
}
|
||||
|
||||
c.maybeAdjustWindowSize(now)
|
||||
c.receiveWindow = c.bytesRead + c.receiveWindowSize
|
||||
return c.receiveWindow
|
||||
}
|
||||
|
||||
// maybeAdjustWindowSize increases the receiveWindowSize if we're sending updates too often.
|
||||
// For details about auto-tuning, see https://docs.google.com/document/d/1SExkMmGiz8VYzV3s9E35JQlJ73vhzCekKkDi85F1qCE/edit?usp=sharing.
|
||||
func (c *receiveFlowController) maybeAdjustWindowSize(now monotime.Time) {
|
||||
bytesReadInEpoch := c.bytesRead - c.epochStartOffset
|
||||
// don't do anything if less than half the window has been consumed
|
||||
if bytesReadInEpoch <= c.receiveWindowSize/2 {
|
||||
return
|
||||
}
|
||||
rtt := c.rttStats.SmoothedRTT()
|
||||
if rtt == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
fraction := float64(bytesReadInEpoch) / float64(c.receiveWindowSize)
|
||||
if now.Sub(c.epochStartTime) < time.Duration(4*fraction*float64(rtt)) {
|
||||
// window is consumed too fast, try to increase the window size
|
||||
newSize := min(2*c.receiveWindowSize, c.maxReceiveWindowSize)
|
||||
if newSize > c.receiveWindowSize && (c.allowWindowIncrease == nil || c.allowWindowIncrease(newSize-c.receiveWindowSize)) {
|
||||
c.receiveWindowSize = newSize
|
||||
}
|
||||
}
|
||||
c.startNewAutoTuningEpoch(now)
|
||||
}
|
||||
|
||||
func (c *receiveFlowController) startNewAutoTuningEpoch(now monotime.Time) {
|
||||
c.epochStartTime = now
|
||||
c.epochStartOffset = c.bytesRead
|
||||
}
|
||||
|
||||
func (c *receiveFlowController) checkFlowControlViolation() bool {
|
||||
return c.highestReceived > c.receiveWindow
|
||||
}
|
||||
+186
@@ -0,0 +1,186 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/monotime"
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/internal/qerr"
|
||||
"github.com/quic-go/quic-go/internal/utils"
|
||||
)
|
||||
|
||||
type connectionFlowController struct {
|
||||
receiveFlowController
|
||||
|
||||
// Protects send-side state, which TryWriteAll can access from application goroutines.
|
||||
sendMutex sync.Mutex
|
||||
bytesSent protocol.ByteCount
|
||||
sendWindow protocol.ByteCount
|
||||
lastBlockedAt protocol.ByteCount
|
||||
}
|
||||
|
||||
// newConnectionFlowController gets a new flow controller for the connection.
|
||||
// It is created before we receive the peer's transport parameters, thus it starts with a sendWindow of 0.
|
||||
func newConnectionFlowController(
|
||||
receiveWindow protocol.ByteCount,
|
||||
maxReceiveWindow protocol.ByteCount,
|
||||
allowWindowIncrease func(size protocol.ByteCount) bool,
|
||||
rttStats *utils.RTTStats,
|
||||
logger utils.Logger,
|
||||
) *connectionFlowController {
|
||||
return &connectionFlowController{
|
||||
receiveFlowController: receiveFlowController{
|
||||
rttStats: rttStats,
|
||||
receiveWindow: receiveWindow,
|
||||
receiveWindowSize: receiveWindow,
|
||||
maxReceiveWindowSize: maxReceiveWindow,
|
||||
allowWindowIncrease: allowWindowIncrease,
|
||||
logger: logger,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// IncrementHighestReceived adds an increment to the highestReceived value
|
||||
func (c *connectionFlowController) IncrementHighestReceived(increment protocol.ByteCount, now monotime.Time) error {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
// If this is the first frame received on this connection, start flow-control auto-tuning.
|
||||
if c.highestReceived == 0 {
|
||||
c.startNewAutoTuningEpoch(now)
|
||||
}
|
||||
c.highestReceived += increment
|
||||
|
||||
if c.checkFlowControlViolation() {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.FlowControlError,
|
||||
ErrorMessage: fmt.Sprintf("received %d bytes for the connection, allowed %d bytes", c.highestReceived, c.receiveWindow),
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *connectionFlowController) AddBytesRead(n protocol.ByteCount) (hasWindowUpdate bool) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
c.addBytesRead(n)
|
||||
return c.hasWindowUpdate()
|
||||
}
|
||||
|
||||
// TryAddBytesSent adds n bytes if sufficient connection-level send credit is available.
|
||||
func (c *connectionFlowController) TryAddBytesSent(n protocol.ByteCount) bool {
|
||||
c.sendMutex.Lock()
|
||||
defer c.sendMutex.Unlock()
|
||||
|
||||
if c.bytesSent > c.sendWindow || n > c.sendWindow-c.bytesSent {
|
||||
return false
|
||||
}
|
||||
c.bytesSent += n
|
||||
return true
|
||||
}
|
||||
|
||||
// AddBytesSentWithLimiter adds the limiter-approved portion of the available connection-level send credit.
|
||||
func (c *connectionFlowController) AddBytesSentWithLimiter(
|
||||
n protocol.ByteCount,
|
||||
limiter func(int) int,
|
||||
) (protocol.ByteCount, bool) {
|
||||
c.sendMutex.Lock()
|
||||
defer c.sendMutex.Unlock()
|
||||
|
||||
if c.bytesSent >= c.sendWindow {
|
||||
return 0, false
|
||||
}
|
||||
n = min(n, c.sendWindow-c.bytesSent)
|
||||
added := min(
|
||||
max(protocol.ByteCount(limiter(int(n))), 0),
|
||||
n,
|
||||
)
|
||||
c.bytesSent += added
|
||||
return added, added < n
|
||||
}
|
||||
|
||||
// UpdateSendWindow is called after receiving a MAX_DATA frame.
|
||||
func (c *connectionFlowController) UpdateSendWindow(offset protocol.ByteCount) (updated bool) {
|
||||
c.sendMutex.Lock()
|
||||
defer c.sendMutex.Unlock()
|
||||
|
||||
if offset > c.sendWindow {
|
||||
c.sendWindow = offset
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (c *connectionFlowController) SendWindowSize() protocol.ByteCount {
|
||||
c.sendMutex.Lock()
|
||||
defer c.sendMutex.Unlock()
|
||||
|
||||
return c.sendWindow - c.bytesSent
|
||||
}
|
||||
|
||||
// IsNewlyBlocked says if it is newly blocked by connection flow control.
|
||||
// For every offset, it only returns true once.
|
||||
// If it is blocked, the offset is returned.
|
||||
func (c *connectionFlowController) IsNewlyBlocked() (bool, protocol.ByteCount) {
|
||||
c.sendMutex.Lock()
|
||||
defer c.sendMutex.Unlock()
|
||||
|
||||
if c.bytesSent < c.sendWindow || c.sendWindow == c.lastBlockedAt {
|
||||
return false, 0
|
||||
}
|
||||
c.lastBlockedAt = c.sendWindow
|
||||
return true, c.sendWindow
|
||||
}
|
||||
|
||||
func (c *connectionFlowController) GetWindowUpdate(now monotime.Time) protocol.ByteCount {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
oldWindowSize := c.receiveWindowSize
|
||||
offset := c.getWindowUpdate(now)
|
||||
if c.logger.Debug() && oldWindowSize < c.receiveWindowSize {
|
||||
c.logger.Debugf("Increasing receive flow control window for the connection to %d kB", c.receiveWindowSize/(1<<10))
|
||||
}
|
||||
return offset
|
||||
}
|
||||
|
||||
// EnsureMinimumWindowSize sets a minimum window size
|
||||
// it should make sure that the connection-level window is increased when a stream-level window grows
|
||||
func (c *connectionFlowController) EnsureMinimumWindowSize(inc protocol.ByteCount, now monotime.Time) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
if inc <= c.receiveWindowSize {
|
||||
return
|
||||
}
|
||||
newSize := min(inc, c.maxReceiveWindowSize)
|
||||
if delta := newSize - c.receiveWindowSize; delta > 0 && c.allowWindowIncrease(delta) {
|
||||
c.receiveWindowSize = newSize
|
||||
if c.logger.Debug() {
|
||||
c.logger.Debugf("Increasing receive flow control window for the connection to %d, in response to stream flow control window increase", newSize)
|
||||
}
|
||||
}
|
||||
c.startNewAutoTuningEpoch(now)
|
||||
}
|
||||
|
||||
// Reset rests the flow controller. This happens when 0-RTT is rejected.
|
||||
// All stream data is invalidated, it's as if we had never opened a stream and never sent any data.
|
||||
// At that point, we only have sent stream data, but we didn't have the keys to open 1-RTT keys yet.
|
||||
func (c *connectionFlowController) Reset() error {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
if c.bytesRead > 0 || c.highestReceived > 0 || !c.epochStartTime.IsZero() {
|
||||
return errors.New("flow controller reset after reading data")
|
||||
}
|
||||
c.sendMutex.Lock()
|
||||
defer c.sendMutex.Unlock()
|
||||
|
||||
c.bytesSent = 0
|
||||
c.lastBlockedAt = 0
|
||||
c.sendWindow = 0
|
||||
return nil
|
||||
}
|
||||
+196
@@ -0,0 +1,196 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/monotime"
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/internal/qerr"
|
||||
"github.com/quic-go/quic-go/internal/utils"
|
||||
)
|
||||
|
||||
type streamFlowController struct {
|
||||
receiveFlowController
|
||||
|
||||
bytesSent protocol.ByteCount
|
||||
sendWindow protocol.ByteCount
|
||||
lastBlockedAt protocol.ByteCount
|
||||
|
||||
streamID protocol.StreamID
|
||||
|
||||
connection *connectionFlowController
|
||||
|
||||
receivedFinalOffset bool
|
||||
}
|
||||
|
||||
// newStreamFlowController gets a new flow controller for a stream.
|
||||
func newStreamFlowController(
|
||||
streamID protocol.StreamID,
|
||||
cfc *connectionFlowController,
|
||||
receiveWindow protocol.ByteCount,
|
||||
maxReceiveWindow protocol.ByteCount,
|
||||
initialSendWindow protocol.ByteCount,
|
||||
rttStats *utils.RTTStats,
|
||||
logger utils.Logger,
|
||||
) *streamFlowController {
|
||||
return &streamFlowController{
|
||||
streamID: streamID,
|
||||
connection: cfc,
|
||||
sendWindow: initialSendWindow,
|
||||
receiveFlowController: receiveFlowController{
|
||||
rttStats: rttStats,
|
||||
receiveWindow: receiveWindow,
|
||||
receiveWindowSize: receiveWindow,
|
||||
maxReceiveWindowSize: maxReceiveWindow,
|
||||
logger: logger,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// UpdateHighestReceived updates the highestReceived value, if the offset is higher.
|
||||
func (c *streamFlowController) UpdateHighestReceived(offset protocol.ByteCount, final bool, now monotime.Time) error {
|
||||
// If the final offset for this stream is already known, check for consistency.
|
||||
if c.receivedFinalOffset {
|
||||
// If we receive another final offset, check that it's the same.
|
||||
if final && offset != c.highestReceived {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.FinalSizeError,
|
||||
ErrorMessage: fmt.Sprintf("received inconsistent final offset for stream %d (old: %d, new: %d bytes)", c.streamID, c.highestReceived, offset),
|
||||
}
|
||||
}
|
||||
// Check that the offset is below the final offset.
|
||||
if offset > c.highestReceived {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.FinalSizeError,
|
||||
ErrorMessage: fmt.Sprintf("received offset %d for stream %d, but final offset was already received at %d", offset, c.streamID, c.highestReceived),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if final {
|
||||
c.receivedFinalOffset = true
|
||||
}
|
||||
if offset == c.highestReceived {
|
||||
return nil
|
||||
}
|
||||
// A higher offset was received before. This can happen due to reordering.
|
||||
if offset < c.highestReceived {
|
||||
if final {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.FinalSizeError,
|
||||
ErrorMessage: fmt.Sprintf("received final offset %d for stream %d, but already received offset %d before", offset, c.streamID, c.highestReceived),
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// If this is the first frame received for this stream, start flow-control auto-tuning.
|
||||
if c.highestReceived == 0 {
|
||||
c.startNewAutoTuningEpoch(now)
|
||||
}
|
||||
increment := offset - c.highestReceived
|
||||
c.highestReceived = offset
|
||||
|
||||
if c.checkFlowControlViolation() {
|
||||
return &qerr.TransportError{
|
||||
ErrorCode: qerr.FlowControlError,
|
||||
ErrorMessage: fmt.Sprintf("received %d bytes on stream %d, allowed %d bytes", offset, c.streamID, c.receiveWindow),
|
||||
}
|
||||
}
|
||||
return c.connection.IncrementHighestReceived(increment, now)
|
||||
}
|
||||
|
||||
func (c *streamFlowController) AddBytesRead(n protocol.ByteCount) (hasStreamWindowUpdate, hasConnWindowUpdate bool) {
|
||||
c.mutex.Lock()
|
||||
c.addBytesRead(n)
|
||||
hasStreamWindowUpdate = c.shouldQueueWindowUpdate()
|
||||
c.mutex.Unlock()
|
||||
hasConnWindowUpdate = c.connection.AddBytesRead(n)
|
||||
return
|
||||
}
|
||||
|
||||
func (c *streamFlowController) Abandon() {
|
||||
c.mutex.Lock()
|
||||
unread := c.highestReceived - c.bytesRead
|
||||
c.bytesRead = c.highestReceived
|
||||
c.mutex.Unlock()
|
||||
if unread > 0 {
|
||||
c.connection.AddBytesRead(unread)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *streamFlowController) UpdateSendWindow(offset protocol.ByteCount) (updated bool) {
|
||||
if offset > c.sendWindow {
|
||||
c.sendWindow = offset
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// TryAddBytesSent adds n bytes if sufficient stream- and connection-level send credit is available.
|
||||
func (c *streamFlowController) TryAddBytesSent(n protocol.ByteCount) bool {
|
||||
if c.bytesSent > c.sendWindow || n > c.sendWindow-c.bytesSent {
|
||||
return false
|
||||
}
|
||||
if !c.connection.TryAddBytesSent(n) {
|
||||
return false
|
||||
}
|
||||
c.bytesSent += n
|
||||
return true
|
||||
}
|
||||
|
||||
// AddBytesSentWithLimiter adds the limiter-approved portion of the available stream- and connection-level send credit.
|
||||
func (c *streamFlowController) AddBytesSentWithLimiter(
|
||||
n protocol.ByteCount,
|
||||
limiter func(int) int,
|
||||
) (protocol.ByteCount, bool) {
|
||||
if c.bytesSent >= c.sendWindow {
|
||||
return 0, false
|
||||
}
|
||||
n = min(n, c.sendWindow-c.bytesSent)
|
||||
added, limited := c.connection.AddBytesSentWithLimiter(n, limiter)
|
||||
c.bytesSent += added
|
||||
return added, limited
|
||||
}
|
||||
|
||||
func (c *streamFlowController) SendWindowSize() protocol.ByteCount {
|
||||
return min(c.sendWindow-c.bytesSent, c.connection.SendWindowSize())
|
||||
}
|
||||
|
||||
func (c *streamFlowController) IsNewlyBlocked() bool {
|
||||
blocked, _ := c.isNewlyBlocked()
|
||||
return blocked
|
||||
}
|
||||
|
||||
func (c *streamFlowController) isNewlyBlocked() (bool, protocol.ByteCount) {
|
||||
if c.bytesSent < c.sendWindow || c.sendWindow == c.lastBlockedAt {
|
||||
return false, 0
|
||||
}
|
||||
c.lastBlockedAt = c.sendWindow
|
||||
return true, c.sendWindow
|
||||
}
|
||||
|
||||
func (c *streamFlowController) shouldQueueWindowUpdate() bool {
|
||||
return !c.receivedFinalOffset && c.hasWindowUpdate()
|
||||
}
|
||||
|
||||
func (c *streamFlowController) GetWindowUpdate(now monotime.Time) protocol.ByteCount {
|
||||
// If we already received the final offset for this stream, the peer won't need any additional flow control credit.
|
||||
if c.receivedFinalOffset {
|
||||
return 0
|
||||
}
|
||||
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
oldWindowSize := c.receiveWindowSize
|
||||
offset := c.getWindowUpdate(now)
|
||||
if c.receiveWindowSize > oldWindowSize { // auto-tuning enlarged the window size
|
||||
c.logger.Debugf("Increasing receive flow control window for stream %d to %d", c.streamID, c.receiveWindowSize)
|
||||
c.connection.EnsureMinimumWindowSize(
|
||||
protocol.ByteCount(float64(c.receiveWindowSize)*protocol.ConnectionFlowControlMultiplier),
|
||||
now,
|
||||
)
|
||||
}
|
||||
return offset
|
||||
}
|
||||
+274
@@ -0,0 +1,274 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"sync"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
list "github.com/quic-go/quic-go/internal/utils/linkedlist"
|
||||
)
|
||||
|
||||
// byteInterval is an interval from one ByteCount to the other
|
||||
type byteInterval struct {
|
||||
Start protocol.ByteCount
|
||||
End protocol.ByteCount
|
||||
}
|
||||
|
||||
var byteIntervalElementPool sync.Pool
|
||||
|
||||
func init() {
|
||||
byteIntervalElementPool = *list.NewPool[byteInterval]()
|
||||
}
|
||||
|
||||
type frameSorterEntry struct {
|
||||
Data []byte
|
||||
DoneCb func()
|
||||
}
|
||||
|
||||
type frameSorter struct {
|
||||
queue map[protocol.ByteCount]frameSorterEntry
|
||||
readPos protocol.ByteCount
|
||||
gaps *list.List[byteInterval]
|
||||
}
|
||||
|
||||
var errDuplicateStreamData = errors.New("duplicate stream data")
|
||||
|
||||
func newFrameSorter() *frameSorter {
|
||||
s := frameSorter{
|
||||
gaps: list.NewWithPool[byteInterval](&byteIntervalElementPool),
|
||||
queue: make(map[protocol.ByteCount]frameSorterEntry),
|
||||
}
|
||||
s.gaps.PushFront(byteInterval{Start: 0, End: protocol.MaxByteCount})
|
||||
return &s
|
||||
}
|
||||
|
||||
func (s *frameSorter) Push(data []byte, offset protocol.ByteCount, doneCb func()) error {
|
||||
err := s.push(data, offset, doneCb)
|
||||
if err == errDuplicateStreamData {
|
||||
if doneCb != nil {
|
||||
doneCb()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *frameSorter) push(data []byte, offset protocol.ByteCount, doneCb func()) error {
|
||||
if len(data) == 0 {
|
||||
return errDuplicateStreamData
|
||||
}
|
||||
|
||||
start := offset
|
||||
end := offset + protocol.ByteCount(len(data))
|
||||
|
||||
if end <= s.gaps.Front().Value.Start {
|
||||
return errDuplicateStreamData
|
||||
}
|
||||
|
||||
startGap, startsInGap := s.findStartGap(start)
|
||||
endGap, endsInGap := s.findEndGap(startGap, end)
|
||||
|
||||
startGapEqualsEndGap := startGap == endGap
|
||||
|
||||
if (startGapEqualsEndGap && end <= startGap.Value.Start) ||
|
||||
(!startGapEqualsEndGap && startGap.Value.End >= endGap.Value.Start && end <= startGap.Value.Start) {
|
||||
return errDuplicateStreamData
|
||||
}
|
||||
|
||||
startGapNext := startGap.Next()
|
||||
startGapEnd := startGap.Value.End // save it, in case startGap is modified
|
||||
endGapStart := endGap.Value.Start // save it, in case endGap is modified
|
||||
endGapEnd := endGap.Value.End // save it, in case endGap is modified
|
||||
var adjustedStartGapEnd bool
|
||||
var wasCut bool
|
||||
|
||||
pos := start
|
||||
var hasReplacedAtLeastOne bool
|
||||
for {
|
||||
oldEntry, ok := s.queue[pos]
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
oldEntryLen := protocol.ByteCount(len(oldEntry.Data))
|
||||
if end-pos > oldEntryLen || (hasReplacedAtLeastOne && end-pos == oldEntryLen) {
|
||||
// The existing frame is shorter than the new frame. Replace it.
|
||||
delete(s.queue, pos)
|
||||
pos += oldEntryLen
|
||||
hasReplacedAtLeastOne = true
|
||||
if oldEntry.DoneCb != nil {
|
||||
oldEntry.DoneCb()
|
||||
}
|
||||
} else {
|
||||
if !hasReplacedAtLeastOne {
|
||||
return errDuplicateStreamData
|
||||
}
|
||||
// The existing frame is longer than the new frame.
|
||||
// Cut the new frame such that the end aligns with the start of the existing frame.
|
||||
data = data[:pos-start]
|
||||
end = pos
|
||||
wasCut = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if !startsInGap && !hasReplacedAtLeastOne {
|
||||
// cut the frame, such that it starts at the start of the gap
|
||||
data = data[startGap.Value.Start-start:]
|
||||
start = startGap.Value.Start
|
||||
wasCut = true
|
||||
}
|
||||
if start <= startGap.Value.Start {
|
||||
if end >= startGap.Value.End {
|
||||
// The frame covers the whole startGap. Delete the gap.
|
||||
s.gaps.Remove(startGap)
|
||||
} else {
|
||||
startGap.Value.Start = end
|
||||
}
|
||||
} else if !hasReplacedAtLeastOne {
|
||||
startGap.Value.End = start
|
||||
adjustedStartGapEnd = true
|
||||
}
|
||||
|
||||
if !startGapEqualsEndGap {
|
||||
s.deleteConsecutive(startGapEnd)
|
||||
var nextGap *list.Element[byteInterval]
|
||||
for gap := startGapNext; gap.Value.End < endGapStart; gap = nextGap {
|
||||
nextGap = gap.Next()
|
||||
s.deleteConsecutive(gap.Value.End)
|
||||
s.gaps.Remove(gap)
|
||||
}
|
||||
}
|
||||
|
||||
if !endsInGap && start != endGapEnd && end > endGapEnd {
|
||||
// cut the frame, such that it ends at the end of the gap
|
||||
data = data[:endGapEnd-start]
|
||||
end = endGapEnd
|
||||
wasCut = true
|
||||
}
|
||||
if end == endGapEnd {
|
||||
if !startGapEqualsEndGap {
|
||||
// The frame covers the whole endGap. Delete the gap.
|
||||
s.gaps.Remove(endGap)
|
||||
}
|
||||
} else {
|
||||
if startGapEqualsEndGap && adjustedStartGapEnd {
|
||||
// The frame split the existing gap into two.
|
||||
s.gaps.InsertAfter(byteInterval{Start: end, End: startGapEnd}, startGap)
|
||||
} else if !startGapEqualsEndGap {
|
||||
endGap.Value.Start = end
|
||||
}
|
||||
}
|
||||
|
||||
if wasCut && len(data) < protocol.MinStreamFrameBufferSize {
|
||||
newData := make([]byte, len(data))
|
||||
copy(newData, data)
|
||||
data = newData
|
||||
if doneCb != nil {
|
||||
doneCb()
|
||||
doneCb = nil
|
||||
}
|
||||
}
|
||||
|
||||
if s.gaps.Len() > protocol.MaxStreamFrameSorterGaps {
|
||||
return errors.New("too many gaps in received data")
|
||||
}
|
||||
|
||||
s.queue[start] = frameSorterEntry{Data: data, DoneCb: doneCb}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *frameSorter) findStartGap(offset protocol.ByteCount) (*list.Element[byteInterval], bool) {
|
||||
for gap := s.gaps.Front(); gap != nil; gap = gap.Next() {
|
||||
if offset >= gap.Value.Start && offset <= gap.Value.End {
|
||||
return gap, true
|
||||
}
|
||||
if offset < gap.Value.Start {
|
||||
return gap, false
|
||||
}
|
||||
}
|
||||
panic("no gap found")
|
||||
}
|
||||
|
||||
func (s *frameSorter) findEndGap(startGap *list.Element[byteInterval], offset protocol.ByteCount) (*list.Element[byteInterval], bool) {
|
||||
for gap := startGap; gap != nil; gap = gap.Next() {
|
||||
if offset >= gap.Value.Start && offset < gap.Value.End {
|
||||
return gap, true
|
||||
}
|
||||
if offset < gap.Value.Start {
|
||||
return gap.Prev(), false
|
||||
}
|
||||
}
|
||||
panic("no gap found")
|
||||
}
|
||||
|
||||
// deleteConsecutive deletes consecutive frames from the queue, starting at pos
|
||||
func (s *frameSorter) deleteConsecutive(pos protocol.ByteCount) {
|
||||
for {
|
||||
oldEntry, ok := s.queue[pos]
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
oldEntryLen := protocol.ByteCount(len(oldEntry.Data))
|
||||
delete(s.queue, pos)
|
||||
if oldEntry.DoneCb != nil {
|
||||
oldEntry.DoneCb()
|
||||
}
|
||||
pos += oldEntryLen
|
||||
}
|
||||
}
|
||||
|
||||
func (s *frameSorter) Pop() (protocol.ByteCount, []byte, func()) {
|
||||
entry, ok := s.queue[s.readPos]
|
||||
if !ok {
|
||||
return s.readPos, nil, nil
|
||||
}
|
||||
delete(s.queue, s.readPos)
|
||||
offset := s.readPos
|
||||
s.readPos += protocol.ByteCount(len(entry.Data))
|
||||
if s.gaps.Front().Value.End <= s.readPos {
|
||||
panic("frame sorter BUG: read position higher than a gap")
|
||||
}
|
||||
return offset, entry.Data, entry.DoneCb
|
||||
}
|
||||
|
||||
// HasMoreData says if there is any more data queued at *any* offset.
|
||||
func (s *frameSorter) HasMoreData() bool {
|
||||
return len(s.queue) > 0
|
||||
}
|
||||
|
||||
var errTooLittleData = errors.New("too little data")
|
||||
|
||||
// Peek copies len(p) consecutive bytes starting at offset into p, without removing them.
|
||||
// It is only possible to peek from an offset where a frame starts.
|
||||
//
|
||||
// If there isn't enough consecutive data available, errTooLittleData is returned.
|
||||
func (s *frameSorter) Peek(offset protocol.ByteCount, p []byte) error {
|
||||
if len(p) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
// first, check if we have enough consecutive data available
|
||||
pos := offset
|
||||
remaining := len(p)
|
||||
for remaining > 0 {
|
||||
entry, ok := s.queue[pos]
|
||||
if !ok {
|
||||
return errTooLittleData
|
||||
}
|
||||
entryLen := len(entry.Data)
|
||||
if remaining <= entryLen {
|
||||
break // enough data available
|
||||
}
|
||||
remaining -= entryLen
|
||||
pos += protocol.ByteCount(entryLen)
|
||||
}
|
||||
|
||||
pos = offset
|
||||
var copied int
|
||||
for copied < len(p) {
|
||||
entry := s.queue[pos] // the entry is guaranteed to exist from the check above
|
||||
copied += copy(p[copied:], entry.Data)
|
||||
pos += protocol.ByteCount(len(entry.Data))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
+295
@@ -0,0 +1,295 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"sync"
|
||||
|
||||
"github.com/quic-go/quic-go/internal/ackhandler"
|
||||
"github.com/quic-go/quic-go/internal/monotime"
|
||||
"github.com/quic-go/quic-go/internal/protocol"
|
||||
"github.com/quic-go/quic-go/internal/utils/ringbuffer"
|
||||
"github.com/quic-go/quic-go/internal/wire"
|
||||
"github.com/quic-go/quic-go/quicvarint"
|
||||
)
|
||||
|
||||
const (
|
||||
maxPathResponses = 256
|
||||
maxControlFrames = 16 << 10
|
||||
)
|
||||
|
||||
// This is the largest possible size of a stream-related control frame
|
||||
// (which is the RESET_STREAM frame).
|
||||
const maxStreamControlFrameSize = 25
|
||||
|
||||
type streamFrameGetter interface {
|
||||
popStreamFrame(protocol.ByteCount, protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame, bool)
|
||||
}
|
||||
|
||||
type streamControlFrameGetter interface {
|
||||
getControlFrame(monotime.Time) (_ ackhandler.Frame, ok, hasMore bool)
|
||||
}
|
||||
|
||||
type framer struct {
|
||||
mutex sync.Mutex
|
||||
|
||||
activeStreams map[protocol.StreamID]streamFrameGetter
|
||||
streamQueue ringbuffer.RingBuffer[protocol.StreamID]
|
||||
streamsWithControlFrames map[protocol.StreamID]streamControlFrameGetter
|
||||
|
||||
controlFrameMutex sync.Mutex
|
||||
controlFrames []wire.Frame
|
||||
pathResponses []*wire.PathResponseFrame
|
||||
connFlowController *connectionFlowController
|
||||
queuedTooManyControlFrames bool
|
||||
}
|
||||
|
||||
func newFramer(connFlowController *connectionFlowController) *framer {
|
||||
return &framer{
|
||||
activeStreams: make(map[protocol.StreamID]streamFrameGetter),
|
||||
streamsWithControlFrames: make(map[protocol.StreamID]streamControlFrameGetter),
|
||||
connFlowController: connFlowController,
|
||||
}
|
||||
}
|
||||
|
||||
func (f *framer) HasData() bool {
|
||||
f.mutex.Lock()
|
||||
hasData := !f.streamQueue.Empty()
|
||||
f.mutex.Unlock()
|
||||
if hasData {
|
||||
return true
|
||||
}
|
||||
f.controlFrameMutex.Lock()
|
||||
defer f.controlFrameMutex.Unlock()
|
||||
return len(f.streamsWithControlFrames) > 0 || len(f.controlFrames) > 0 || len(f.pathResponses) > 0
|
||||
}
|
||||
|
||||
func (f *framer) QueueControlFrame(frame wire.Frame) {
|
||||
f.controlFrameMutex.Lock()
|
||||
defer f.controlFrameMutex.Unlock()
|
||||
|
||||
if pr, ok := frame.(*wire.PathResponseFrame); ok {
|
||||
// Only queue up to maxPathResponses PATH_RESPONSE frames.
|
||||
// This limit should be high enough to never be hit in practice,
|
||||
// unless the peer is doing something malicious.
|
||||
if len(f.pathResponses) >= maxPathResponses {
|
||||
return
|
||||
}
|
||||
f.pathResponses = append(f.pathResponses, pr)
|
||||
return
|
||||
}
|
||||
// This is a hack.
|
||||
if len(f.controlFrames) >= maxControlFrames {
|
||||
f.queuedTooManyControlFrames = true
|
||||
return
|
||||
}
|
||||
f.controlFrames = append(f.controlFrames, frame)
|
||||
}
|
||||
|
||||
func (f *framer) Append(
|
||||
frames []ackhandler.Frame,
|
||||
streamFrames []ackhandler.StreamFrame,
|
||||
maxLen protocol.ByteCount,
|
||||
now monotime.Time,
|
||||
v protocol.Version,
|
||||
) ([]ackhandler.Frame, []ackhandler.StreamFrame, protocol.ByteCount) {
|
||||
f.controlFrameMutex.Lock()
|
||||
frames, controlFrameLen := f.appendControlFrames(frames, maxLen, now, v)
|
||||
maxLen -= controlFrameLen
|
||||
|
||||
var lastFrame ackhandler.StreamFrame
|
||||
var streamFrameLen protocol.ByteCount
|
||||
f.mutex.Lock()
|
||||
// pop STREAM frames, until less than 128 bytes are left in the packet
|
||||
numActiveStreams := f.streamQueue.Len()
|
||||
for range numActiveStreams {
|
||||
if protocol.MinStreamFrameSize > maxLen {
|
||||
break
|
||||
}
|
||||
sf, blocked := f.getNextStreamFrame(maxLen, v)
|
||||
if sf.Frame != nil {
|
||||
streamFrames = append(streamFrames, sf)
|
||||
maxLen -= sf.Frame.Length(v)
|
||||
lastFrame = sf
|
||||
streamFrameLen += sf.Frame.Length(v)
|
||||
}
|
||||
// If the stream just became blocked on stream flow control, attempt to pack the
|
||||
// STREAM_DATA_BLOCKED into the same packet.
|
||||
if blocked != nil {
|
||||
l := blocked.Length(v)
|
||||
// In case it doesn't fit, queue it for the next packet.
|
||||
if maxLen < l {
|
||||
f.controlFrames = append(f.controlFrames, blocked)
|
||||
break
|
||||
}
|
||||
frames = append(frames, ackhandler.Frame{Frame: blocked})
|
||||
maxLen -= l
|
||||
controlFrameLen += l
|
||||
}
|
||||
}
|
||||
|
||||
// The only way to become blocked on connection-level flow control is by sending STREAM frames.
|
||||
if isBlocked, offset := f.connFlowController.IsNewlyBlocked(); isBlocked {
|
||||
blocked := &wire.DataBlockedFrame{MaximumData: offset}
|
||||
l := blocked.Length(v)
|
||||
// In case it doesn't fit, queue it for the next packet.
|
||||
if maxLen >= l {
|
||||
frames = append(frames, ackhandler.Frame{Frame: blocked})
|
||||
controlFrameLen += l
|
||||
} else {
|
||||
f.controlFrames = append(f.controlFrames, blocked)
|
||||
}
|
||||
}
|
||||
|
||||
f.mutex.Unlock()
|
||||
f.controlFrameMutex.Unlock()
|
||||
|
||||
if lastFrame.Frame != nil {
|
||||
// account for the smaller size of the last STREAM frame
|
||||
streamFrameLen -= lastFrame.Frame.Length(v)
|
||||
lastFrame.Frame.DataLenPresent = false
|
||||
streamFrameLen += lastFrame.Frame.Length(v)
|
||||
}
|
||||
|
||||
return frames, streamFrames, controlFrameLen + streamFrameLen
|
||||
}
|
||||
|
||||
func (f *framer) appendControlFrames(
|
||||
frames []ackhandler.Frame,
|
||||
maxLen protocol.ByteCount,
|
||||
now monotime.Time,
|
||||
v protocol.Version,
|
||||
) ([]ackhandler.Frame, protocol.ByteCount) {
|
||||
var length protocol.ByteCount
|
||||
// add a PATH_RESPONSE first, but only pack a single PATH_RESPONSE per packet
|
||||
if len(f.pathResponses) > 0 {
|
||||
frame := f.pathResponses[0]
|
||||
frameLen := frame.Length(v)
|
||||
if frameLen <= maxLen {
|
||||
frames = append(frames, ackhandler.Frame{Frame: frame})
|
||||
length += frameLen
|
||||
f.pathResponses = f.pathResponses[1:]
|
||||
}
|
||||
}
|
||||
|
||||
// add stream-related control frames
|
||||
for id, str := range f.streamsWithControlFrames {
|
||||
start:
|
||||
remainingLen := maxLen - length
|
||||
if remainingLen <= maxStreamControlFrameSize {
|
||||
break
|
||||
}
|
||||
fr, ok, hasMore := str.getControlFrame(now)
|
||||
if !hasMore {
|
||||
delete(f.streamsWithControlFrames, id)
|
||||
}
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
frames = append(frames, fr)
|
||||
length += fr.Frame.Length(v)
|
||||
if hasMore {
|
||||
// It is rare that a stream has more than one control frame to queue.
|
||||
// We don't want to spawn another loop for just to cover that case.
|
||||
goto start
|
||||
}
|
||||
}
|
||||
|
||||
for len(f.controlFrames) > 0 {
|
||||
frame := f.controlFrames[len(f.controlFrames)-1]
|
||||
frameLen := frame.Length(v)
|
||||
if length+frameLen > maxLen {
|
||||
break
|
||||
}
|
||||
frames = append(frames, ackhandler.Frame{Frame: frame})
|
||||
length += frameLen
|
||||
f.controlFrames = f.controlFrames[:len(f.controlFrames)-1]
|
||||
}
|
||||
|
||||
return frames, length
|
||||
}
|
||||
|
||||
// QueuedTooManyControlFrames says if the control frame queue exceeded its maximum queue length.
|
||||
// This is a hack.
|
||||
// It is easier to implement than propagating an error return value in QueueControlFrame.
|
||||
// The correct solution would be to queue frames with their respective structs.
|
||||
// See https://github.com/quic-go/quic-go/issues/4271 for the queueing of stream-related control frames.
|
||||
func (f *framer) QueuedTooManyControlFrames() bool {
|
||||
return f.queuedTooManyControlFrames
|
||||
}
|
||||
|
||||
func (f *framer) AddActiveStream(id protocol.StreamID, str streamFrameGetter) {
|
||||
f.mutex.Lock()
|
||||
if _, ok := f.activeStreams[id]; !ok {
|
||||
f.streamQueue.PushBack(id)
|
||||
f.activeStreams[id] = str
|
||||
}
|
||||
f.mutex.Unlock()
|
||||
}
|
||||
|
||||
func (f *framer) AddStreamWithControlFrames(id protocol.StreamID, str streamControlFrameGetter) {
|
||||
f.controlFrameMutex.Lock()
|
||||
if _, ok := f.streamsWithControlFrames[id]; !ok {
|
||||
f.streamsWithControlFrames[id] = str
|
||||
}
|
||||
f.controlFrameMutex.Unlock()
|
||||
}
|
||||
|
||||
// RemoveActiveStream is called when a stream completes.
|
||||
func (f *framer) RemoveActiveStream(id protocol.StreamID) {
|
||||
f.mutex.Lock()
|
||||
delete(f.activeStreams, id)
|
||||
// We don't delete the stream from the streamQueue,
|
||||
// since we'd have to iterate over the ringbuffer.
|
||||
// Instead, we check if the stream is still in activeStreams when appending STREAM frames.
|
||||
f.mutex.Unlock()
|
||||
}
|
||||
|
||||
func (f *framer) getNextStreamFrame(maxLen protocol.ByteCount, v protocol.Version) (ackhandler.StreamFrame, *wire.StreamDataBlockedFrame) {
|
||||
id := f.streamQueue.PopFront()
|
||||
// This should never return an error. Better check it anyway.
|
||||
// The stream will only be in the streamQueue, if it enqueued itself there.
|
||||
str, ok := f.activeStreams[id]
|
||||
// The stream might have been removed after being enqueued.
|
||||
if !ok {
|
||||
return ackhandler.StreamFrame{}, nil
|
||||
}
|
||||
// For the last STREAM frame, we'll remove the DataLen field later.
|
||||
// Therefore, we can pretend to have more bytes available when popping
|
||||
// the STREAM frame (which will always have the DataLen set).
|
||||
maxLen += protocol.ByteCount(quicvarint.Len(uint64(maxLen)))
|
||||
frame, blocked, hasMoreData := str.popStreamFrame(maxLen, v)
|
||||
if hasMoreData { // put the stream back in the queue (at the end)
|
||||
f.streamQueue.PushBack(id)
|
||||
} else { // no more data to send. Stream is not active
|
||||
delete(f.activeStreams, id)
|
||||
}
|
||||
// Note that the frame.Frame can be nil:
|
||||
// * if the stream was canceled after it said it had data
|
||||
// * the remaining size doesn't allow us to add another STREAM frame
|
||||
return frame, blocked
|
||||
}
|
||||
|
||||
func (f *framer) Handle0RTTRejection() {
|
||||
f.mutex.Lock()
|
||||
defer f.mutex.Unlock()
|
||||
f.controlFrameMutex.Lock()
|
||||
defer f.controlFrameMutex.Unlock()
|
||||
|
||||
f.streamQueue.Clear()
|
||||
for id := range f.activeStreams {
|
||||
delete(f.activeStreams, id)
|
||||
}
|
||||
clear(f.streamsWithControlFrames)
|
||||
var j int
|
||||
for i, frame := range f.controlFrames {
|
||||
switch frame.(type) {
|
||||
case *wire.MaxDataFrame, *wire.MaxStreamDataFrame, *wire.MaxStreamsFrame,
|
||||
*wire.DataBlockedFrame, *wire.StreamDataBlockedFrame, *wire.StreamsBlockedFrame:
|
||||
continue
|
||||
default:
|
||||
f.controlFrames[j] = f.controlFrames[i]
|
||||
j++
|
||||
}
|
||||
}
|
||||
f.controlFrames = slices.Delete(f.controlFrames, j, len(f.controlFrames))
|
||||
}
|
||||
Loaded 100 of 250 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user