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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions docs/webtransport.md
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,8 @@ a terminal.
Use `--wt-mode datagram` for unreliable datagrams. `--wt-datagram-mode lines`
sends one datagram per line; `binary` sends 1 KiB chunks. Received datagrams
are JSON Lines records with `sequence`, `length`, and base64 `data` fields.
If an outgoing datagram is too large, the error reports its size and the
current underlying QUIC limit (before HTTP/3 framing overhead).
Datagram input ending does not close the session. Use Ctrl+C when the peer does
not close it.

Expand Down
128 changes: 97 additions & 31 deletions internal/wt/wt.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (
"io"
"unicode/utf8"

"github.com/quic-go/quic-go"
"github.com/quic-go/webtransport-go"
"github.com/ryanfowler/fetch/internal/core"
)
Expand Down Expand Up @@ -45,7 +46,10 @@ func Run(ctx context.Context, cfg Config) error {
return runStream(ctx, cfg)
}

type streamResult struct{ err error }
type streamResult struct {
err error
initiatedShutdown bool
}

func runStream(ctx context.Context, cfg Config) error {
stream, err := cfg.Session.OpenStream(ctx)
Expand All @@ -71,11 +75,12 @@ func runStream(ctx context.Context, cfg Config) error {
writeDone := make(chan streamResult, 1)
go func() {
err := writeStream(workCtx, stream, cfg)
if err != nil {
initiatedShutdown := err != nil && workCtx.Err() == nil
if initiatedShutdown {
_ = cfg.Session.Close()
cancel()
}
writeDone <- streamResult{err}
writeDone <- streamResult{err: err, initiatedShutdown: initiatedShutdown}
}()

reader := io.Reader(stream)
Expand All @@ -87,11 +92,17 @@ func runStream(ctx context.Context, cfg Config) error {
}
buf := make([]byte, 32*1024)
readErr := error(nil)
outputErr := false
for {
n, read := reader.Read(buf)
if n > 0 {
if _, write := output.Write(buf[:n]); write != nil {
readErr = write
written, write := output.Write(buf[:n])
if write == nil && written != n {
write = io.ErrShortWrite
}
if write != nil {
readErr = fmt.Errorf("write WebTransport stream output: %w", write)
outputErr = true
cancel()
_ = cfg.Session.Close()
break
Expand All @@ -107,33 +118,45 @@ func runStream(ctx context.Context, cfg Config) error {
}
}
if safe != nil && readErr == nil {
readErr = safe.Flush()
if err := safe.Flush(); err != nil {
readErr = fmt.Errorf("write WebTransport stream output: %w", err)
outputErr = true
cancel()
_ = cfg.Session.Close()
}
}
write := <-writeDone
if outputErr {
_ = cfg.Session.Close()
return readErr
}
if write.initiatedShutdown {
return write.err
}
write := (<-writeDone).err
if readErr != nil && !isCleanClose(readErr) {
_ = cfg.Session.Close()
return fmt.Errorf("read WebTransport stream: %w", readErr)
}
if write != nil {
if write.err != nil {
_ = cfg.Session.Close()
return write
return write.err
}
return nil
}

func writeStream(ctx context.Context, stream io.WriteCloser, cfg Config) error {
buf := make([]byte, 32*1024)
write := func(r io.Reader) error {
if r == nil {
return nil
}
_, err := io.CopyBuffer(stream, r, make([]byte, 32*1024))
if err != nil {
if err := copyContext(ctx, stream, r, buf); err != nil {
return fmt.Errorf("write WebTransport stream: %w", err)
}
return nil
}
if cfg.InitialPayloadSet {
if _, err := stream.Write(cfg.InitialPayload); err != nil {
if err := writeFull(ctx, stream, cfg.InitialPayload); err != nil {
return fmt.Errorf("write initial WebTransport payload: %w", err)
}
}
Expand All @@ -145,7 +168,7 @@ func writeStream(ctx context.Context, stream io.WriteCloser, cfg Config) error {
closeInput(cfg.InitialReader)
}
if cfg.Stdin != nil {
if err := writeContext(ctx, stream, cfg.Stdin); err != nil {
if err := write(cfg.Stdin); err != nil {
closeInput(cfg.Stdin)
return err
}
Expand All @@ -157,8 +180,7 @@ func writeStream(ctx context.Context, stream io.WriteCloser, cfg Config) error {
return nil
}

func writeContext(ctx context.Context, dst io.Writer, src io.Reader) error {
buf := make([]byte, 32*1024)
func copyContext(ctx context.Context, dst io.Writer, src io.Reader, buf []byte) error {
for {
select {
case <-ctx.Done():
Expand All @@ -167,7 +189,7 @@ func writeContext(ctx context.Context, dst io.Writer, src io.Reader) error {
}
n, err := src.Read(buf)
if n > 0 {
if _, e := dst.Write(buf[:n]); e != nil {
if e := writeFull(ctx, dst, buf[:n]); e != nil {
return e
}
}
Expand All @@ -180,6 +202,30 @@ func writeContext(ctx context.Context, dst io.Writer, src io.Reader) error {
}
}

func writeFull(ctx context.Context, dst io.Writer, p []byte) error {
for len(p) > 0 {
select {
case <-ctx.Done():
return context.Cause(ctx)
default:
}
n, err := dst.Write(p)
if n < 0 || n > len(p) {
return fmt.Errorf("invalid write result %d for %d-byte buffer", n, len(p))
}
if n > 0 {
p = p[n:]
}
if err != nil {
return err
}
if n == 0 {
return io.ErrNoProgress
}
}
return nil
}

func runDatagrams(ctx context.Context, cfg Config) error {
workCtx, cancel := context.WithCancel(ctx)
defer cancel()
Expand All @@ -205,25 +251,29 @@ func runDatagrams(ctx context.Context, cfg Config) error {

var sendErr error
if cfg.InitialPayloadSet {
sendErr = sendDatagram(cfg.Session, cfg.InitialPayload)
sendErr = sendDatagramContext(workCtx, cfg.Session, cfg.InitialPayload)
}
if sendErr == nil && cfg.InitialReader != nil {
data, err := core.ReadAllLimited(cfg.InitialReader, core.MaxCompositeMaterialization, "WebTransport initial datagram")
closeInput(cfg.InitialReader)
if err != nil {
sendErr = err
} else {
sendErr = sendDatagram(cfg.Session, data)
sendErr = sendDatagramContext(workCtx, cfg.Session, data)
}
}
if sendErr == nil && cfg.Stdin != nil {
sendErr = sendInputDatagrams(workCtx, cfg.Session, cfg.Stdin, cfg.DatagramMode)
closeInput(cfg.Stdin)
}
if sendErr != nil {
receiveEnded := workCtx.Err() != nil
cancel()
_ = cfg.Session.Close()
<-receiveDone
receiveErr := <-receiveDone
if receiveEnded {
return normalizeDatagramClose(ctx, receiveErr)
}
return sendErr
}
return normalizeDatagramClose(ctx, <-receiveDone)
Expand All @@ -242,9 +292,13 @@ func receiveDatagrams(ctx context.Context, cfg Config) error {
Data string `json:"data"`
}{seq, len(data), base64.StdEncoding.EncodeToString(data)})
record = append(record, '\n')
if _, err := cfg.Stdout.Write(record); err != nil {
n, err := cfg.Stdout.Write(record)
if err == nil && n != len(record) {
err = io.ErrShortWrite
}
if err != nil {
_ = cfg.Session.Close()
return err
return fmt.Errorf("write WebTransport datagram output: %w", err)
}
seq++
}
Expand Down Expand Up @@ -278,18 +332,36 @@ func normalizeDatagramClose(parent context.Context, err error) error {

func sendDatagram(s Session, data []byte) error {
if err := s.SendDatagram(data); err != nil {
var sizeErr *quic.DatagramTooLargeError
if errors.As(err, &sizeErr) && sizeErr.MaxDatagramPayloadSize > 0 {
return fmt.Errorf("send WebTransport datagram (%d bytes; current QUIC limit %d bytes before HTTP/3 overhead): %w", len(data), sizeErr.MaxDatagramPayloadSize, err)
}
return fmt.Errorf("send WebTransport datagram (%d bytes): %w", len(data), err)
}
return nil
}

func sendDatagramContext(ctx context.Context, s Session, data []byte) error {
select {
case <-ctx.Done():
return context.Cause(ctx)
default:
return sendDatagram(s, data)
}
}

func sendInputDatagrams(ctx context.Context, s Session, r io.Reader, mode core.WTDatagramMode) error {
if mode == core.WTDatagramBinary {
buf := make([]byte, core.MaxWebTransportBinaryChunk)
for {
select {
case <-ctx.Done():
return context.Cause(ctx)
default:
}
n, err := r.Read(buf)
if n > 0 {
if e := sendDatagram(s, append([]byte(nil), buf[:n]...)); e != nil {
if e := sendDatagramContext(ctx, s, append([]byte(nil), buf[:n]...)); e != nil {
return e
}
}
Expand All @@ -308,7 +380,7 @@ func sendInputDatagrams(ctx context.Context, s Session, r io.Reader, mode core.W
return core.LimitError{Subsystem: "WebTransport datagram line", Limit: core.MaxWebTransportDatagramLine}
}
if len(line) > 0 || err != io.EOF {
if e := sendDatagram(s, line); e != nil {
if e := sendDatagramContext(ctx, s, line); e != nil {
return e
}
}
Expand All @@ -318,11 +390,6 @@ func sendInputDatagrams(ctx context.Context, s Session, r io.Reader, mode core.W
if err != nil {
return fmt.Errorf("read WebTransport datagram input: %w", err)
}
select {
case <-ctx.Done():
return context.Cause(ctx)
default:
}
}
}

Expand Down Expand Up @@ -366,7 +433,7 @@ func (w *terminalWriter) Write(p []byte) (int, error) {
return len(p), nil
}
out := core.AppendTerminalSafeBytes(nil, w.pending[:cut])
if _, err := w.dst.Write(out); err != nil {
if err := writeFull(context.Background(), w.dst, out); err != nil {
return 0, err
}
w.pending = append(w.pending[:0], w.pending[cut:]...)
Expand All @@ -391,6 +458,5 @@ func (w *terminalWriter) Flush() error {
}
out := core.AppendTerminalSafeBytes(nil, w.pending)
w.pending = nil
_, err := w.dst.Write(out)
return err
return writeFull(context.Background(), w.dst, out)
}
Loading
Loading