vendor backend

Signed-off-by: RonniSkansing <rskansing@gmail.com>
This commit is contained in:
RonniSkansing committed 2026-09-16 23:17:35 +02:00
1 parent 98ad0bdf4a
commit 2028391c3c
250 files changed
+47469

No files matched your search

+128
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
File diff suppressed because it is too large. Load diff
+98
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
File diff suppressed because it is too large. Load diff
+1152
View File
File diff suppressed because it is too large. Load diff
+147
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -0,0 +1,9 @@
# HTTP/3
[![Documentation](https://img.shields.io/badge/docs-quic--go.net-red?style=flat)](https://quic-go.net/docs/)
[![PkgGoDev](https://pkg.go.dev/badge/github.com/quic-go/quic-go/http3)](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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
+20
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -0,0 +1,59 @@
# Fuzzing
[![Documentation](https://img.shields.io/badge/OSS--Fuzz-Introspector-red?style=flat)](https://introspector.oss-fuzz.com/project-profile?project=quic-go)
[![ClusterFuzz coverage](https://img.shields.io/codecov/c/github/quic-go/quic-go/master.svg?flag=clusterfuzz&label=ClusterFuzz%20coverage&logo=codecov&logoColor=white&style=flat)](https://app.codecov.io/gh/quic-go/quic-go?flags%5B0%5D=clusterfuzz)
[![ClusterFuzz Lite Batch coverage](https://img.shields.io/codecov/c/github/quic-go/quic-go/master.svg?flag=clusterfuzz-lite-batch&label=ClusterFuzz%20Lite%20Batch%20coverage&logo=codecov&logoColor=white&style=flat)](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
View File
@@ -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
View File
@@ -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
[![Documentation](https://img.shields.io/badge/docs-quic--go.net-red?style=flat)](https://quic-go.net/docs/)
[![PkgGoDev](https://pkg.go.dev/badge/github.com/quic-go/quic-go)](https://pkg.go.dev/github.com/quic-go/quic-go)
[![Code Coverage](https://img.shields.io/codecov/c/github/quic-go/quic-go/master.svg?style=flat-square)](https://codecov.io/gh/quic-go/quic-go/)
[![Fuzzing Status](https://oss-fuzz-build-logs.storage.googleapis.com/badges/quic-go.svg)](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. | ![GitHub Repo stars](https://img.shields.io/github/stars/AdguardTeam/AdGuardHome?style=flat-square) |
| [algernon](https://github.com/xyproto/algernon) | Small self-contained pure-Go web server with Lua, Markdown, HTTP/2, QUIC, Redis and PostgreSQL support | ![GitHub Repo stars](https://img.shields.io/github/stars/xyproto/algernon?style=flat-square) |
| [caddy](https://github.com/caddyserver/caddy/) | Fast, multi-platform web server with automatic HTTPS | ![GitHub Repo stars](https://img.shields.io/github/stars/caddyserver/caddy?style=flat-square) |
| [cloudflared](https://github.com/cloudflare/cloudflared) | A tunneling daemon that proxies traffic from the Cloudflare network to your origins | ![GitHub Repo stars](https://img.shields.io/github/stars/cloudflare/cloudflared?style=flat-square) |
| [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 | ![GitHub Repo stars](https://img.shields.io/github/stars/fatedier/frp?style=flat-square) |
| [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 | ![GitHub Repo stars](https://img.shields.io/github/stars/libp2p/go-libp2p?style=flat-square) |
| [gost](https://github.com/go-gost/gost) | A simple security tunnel written in Go | ![GitHub Repo stars](https://img.shields.io/github/stars/go-gost/gost?style=flat-square) |
| [Hysteria](https://github.com/apernet/hysteria) | A powerful, lightning fast and censorship resistant proxy | ![GitHub Repo stars](https://img.shields.io/github/stars/apernet/hysteria?style=flat-square) |
| [Mercure](https://github.com/dunglas/mercure) | An open, easy, fast, reliable and battery-efficient solution for real-time communications | ![GitHub Repo stars](https://img.shields.io/github/stars/dunglas/mercure?style=flat-square) |
| [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. | ![GitHub Repo stars](https://img.shields.io/github/stars/NodePassProject/nodepass?style=flat-square) |
| [OONI Probe](https://github.com/ooni/probe-cli) | Next generation OONI Probe. Library and CLI tool. | ![GitHub Repo stars](https://img.shields.io/github/stars/ooni/probe-cli?style=flat-square) |
| [reverst](https://github.com/flipt-io/reverst) | Reverse Tunnels in Go over HTTP/3 and QUIC | ![GitHub Repo stars](https://img.shields.io/github/stars/flipt-io/reverst?style=flat-square) |
| [RoadRunner](https://github.com/roadrunner-server/roadrunner) | High-performance PHP application server, process manager written in Go and powered with plugins | ![GitHub Repo stars](https://img.shields.io/github/stars/roadrunner-server/roadrunner?style=flat-square) |
| [syncthing](https://github.com/syncthing/syncthing/) | Open Source Continuous File Synchronization | ![GitHub Repo stars](https://img.shields.io/github/stars/syncthing/syncthing?style=flat-square) |
| [traefik](https://github.com/traefik/traefik) | The Cloud Native Application Proxy | ![GitHub Repo stars](https://img.shields.io/github/stars/traefik/traefik?style=flat-square) |
| [v2ray-core](https://github.com/v2fly/v2ray-core) | A platform for building proxies to bypass network restrictions | ![GitHub Repo stars](https://img.shields.io/github/stars/v2fly/v2ray-core?style=flat-square) |
| [YoMo](https://github.com/yomorun/yomo) | Streaming Serverless Framework for Geo-distributed System | ![GitHub Repo stars](https://img.shields.io/github/stars/yomorun/yomo?style=flat-square) |
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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 &copy
}
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
View File
@@ -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
View File
@@ -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")
}
}
File diff suppressed because it is too large. Load diff
+315
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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