From 38eb79b7cffbfc6c205c85a687c609a0953bed83 Mon Sep 17 00:00:00 2001 From: Ryan Fowler Date: Sun, 30 Aug 2026 20:50:35 +0000 Subject: [PATCH] fix(webtransport): harden session I/O --- docs/webtransport.md | 2 + internal/wt/wt.go | 128 +++++++++---- internal/wt/wt_test.go | 237 +++++++++++++++++++++++- skills/fetch/references/webtransport.md | 1 + 4 files changed, 329 insertions(+), 39 deletions(-) diff --git a/docs/webtransport.md b/docs/webtransport.md index 95c16e75..b1514733 100644 --- a/docs/webtransport.md +++ b/docs/webtransport.md @@ -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. diff --git a/internal/wt/wt.go b/internal/wt/wt.go index 0c3ab80b..93fccdd3 100644 --- a/internal/wt/wt.go +++ b/internal/wt/wt.go @@ -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" ) @@ -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) @@ -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) @@ -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 @@ -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) } } @@ -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 } @@ -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(): @@ -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 } } @@ -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() @@ -205,7 +251,7 @@ 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") @@ -213,7 +259,7 @@ func runDatagrams(ctx context.Context, cfg Config) error { if err != nil { sendErr = err } else { - sendErr = sendDatagram(cfg.Session, data) + sendErr = sendDatagramContext(workCtx, cfg.Session, data) } } if sendErr == nil && cfg.Stdin != nil { @@ -221,9 +267,13 @@ func runDatagrams(ctx context.Context, cfg Config) error { 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) @@ -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++ } @@ -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 } } @@ -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 } } @@ -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: - } } } @@ -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:]...) @@ -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) } diff --git a/internal/wt/wt_test.go b/internal/wt/wt_test.go index 6391e7f8..0670f194 100644 --- a/internal/wt/wt_test.go +++ b/internal/wt/wt_test.go @@ -6,23 +6,45 @@ import ( "errors" "io" "strings" + "sync" "testing" + "time" + "github.com/quic-go/quic-go" "github.com/ryanfowler/fetch/internal/core" ) type fakeSession struct { - stream *fakeStream - datagrams [][]byte - received [][]byte + stream *fakeStream + datagrams [][]byte + received [][]byte + sendErr error + receiveGate <-chan struct{} + sendReady chan struct{} + sendTarget int + sendOnce sync.Once } func (s *fakeSession) OpenStream(context.Context) (io.ReadWriteCloser, error) { return s.stream, nil } func (s *fakeSession) SendDatagram(p []byte) error { + if s.sendErr != nil { + return s.sendErr + } s.datagrams = append(s.datagrams, append([]byte(nil), p...)) + if s.sendReady != nil && len(s.datagrams) == s.sendTarget { + s.sendOnce.Do(func() { close(s.sendReady) }) + } return nil } -func (s *fakeSession) ReceiveDatagram(context.Context) ([]byte, error) { +func (s *fakeSession) ReceiveDatagram(ctx context.Context) ([]byte, error) { + if s.receiveGate != nil { + select { + case <-s.receiveGate: + case <-ctx.Done(): + return nil, context.Cause(ctx) + } + s.receiveGate = nil + } if len(s.received) == 0 { return nil, errors.New("done") } @@ -34,8 +56,9 @@ func (s *fakeSession) Close() error { return nil } type fakeStream struct { bytes.Buffer - input []byte - closed bool + input []byte + closed bool + maxWrite int } func (s *fakeStream) Read(p []byte) (int, error) { @@ -46,6 +69,12 @@ func (s *fakeStream) Read(p []byte) (int, error) { s.input = s.input[n:] return n, nil } +func (s *fakeStream) Write(p []byte) (int, error) { + if s.maxWrite > 0 && len(p) > s.maxWrite { + p = p[:s.maxWrite] + } + return s.Buffer.Write(p) +} func (s *fakeStream) Close() error { s.closed = true; return nil } func TestRunStreamDefersInputAndReadsAfterEOF(t *testing.T) { @@ -67,9 +96,12 @@ func TestRunStreamDefersInputAndReadsAfterEOF(t *testing.T) { } func TestRunDatagramsUsesOneInitialPayloadAndJSONLines(t *testing.T) { - s := &fakeSession{received: [][]byte{{0, 1}, {}}} + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + ready := make(chan struct{}) + s := &fakeSession{received: [][]byte{{0, 1}, {}}, receiveGate: ready, sendReady: ready, sendTarget: 3} var out bytes.Buffer - err := Run(context.Background(), Config{Session: s, Stdout: &out, Mode: core.WTDatagram, InitialReader: strings.NewReader("a\nb"), Stdin: strings.NewReader("one\ntwo\n")}) + err := Run(ctx, Config{Session: s, Stdout: &out, Mode: core.WTDatagram, InitialReader: strings.NewReader("a\nb"), Stdin: strings.NewReader("one\ntwo\n")}) if err == nil { t.Fatal("expected receive loop error") } @@ -82,6 +114,195 @@ func TestRunDatagramsUsesOneInitialPayloadAndJSONLines(t *testing.T) { } } +func TestRunStreamCompletesShortWrites(t *testing.T) { + s := &fakeSession{stream: &fakeStream{maxWrite: 2}} + var out bytes.Buffer + err := Run(context.Background(), Config{Session: s, Stdout: &out, InitialPayloadSet: true, InitialPayload: []byte("first"), InitialReader: strings.NewReader("second")}) + if err != nil { + t.Fatal(err) + } + if got := s.stream.String(); got != "firstsecond" { + t.Fatalf("sent %q", got) + } +} + +func TestRunStreamRejectsShortOutputWrite(t *testing.T) { + s := &fakeSession{stream: &fakeStream{input: []byte("reply")}} + err := Run(context.Background(), Config{Session: s, Stdout: &shortWriter{}}) + if !errors.Is(err, io.ErrShortWrite) || !strings.Contains(err.Error(), "stream output") { + t.Fatalf("got %v", err) + } +} + +type writeFailureStream struct { + closed <-chan struct{} + err error +} + +func (s *writeFailureStream) Read([]byte) (int, error) { + <-s.closed + return 0, s.err +} + +func (s *writeFailureStream) Write([]byte) (int, error) { return 0, nil } +func (s *writeFailureStream) Close() error { return nil } + +type writeFailureSession struct { + stream *writeFailureStream + closed chan struct{} + once sync.Once +} + +func (s *writeFailureSession) OpenStream(context.Context) (io.ReadWriteCloser, error) { + return s.stream, nil +} +func (s *writeFailureSession) SendDatagram([]byte) error { return nil } +func (s *writeFailureSession) ReceiveDatagram(context.Context) ([]byte, error) { + return nil, errors.New("unused") +} +func (s *writeFailureSession) Close() error { + s.once.Do(func() { close(s.closed) }) + return nil +} + +func TestRunStreamPreservesInitiatingWriteError(t *testing.T) { + closed := make(chan struct{}) + readErr := errors.New("session closed after write failure") + s := &writeFailureSession{closed: closed, stream: &writeFailureStream{closed: closed, err: readErr}} + err := Run(context.Background(), Config{Session: s, Stdout: io.Discard, InitialPayloadSet: true, InitialPayload: []byte("payload")}) + if !errors.Is(err, io.ErrNoProgress) || errors.Is(err, readErr) { + t.Fatalf("got %v", err) + } +} + +type shortWriter struct{ bytes.Buffer } + +func (w *shortWriter) Write(p []byte) (int, error) { + if len(p) == 0 { + return 0, nil + } + return w.Buffer.Write(p[:len(p)-1]) +} + +type blockingReadCloser struct { + closed chan struct{} + once sync.Once +} + +func (r *blockingReadCloser) Read([]byte) (int, error) { + <-r.closed + return 0, errors.New("input closed") +} + +func (r *blockingReadCloser) Close() error { + r.once.Do(func() { close(r.closed) }) + return nil +} + +func TestRunStreamFlushErrorCancelsBlockingInput(t *testing.T) { + s := &fakeSession{stream: &fakeStream{input: []byte{0xe2}}} + stdin := &blockingReadCloser{closed: make(chan struct{})} + done := make(chan error, 1) + go func() { + done <- Run(context.Background(), Config{Session: s, Stdout: &shortWriter{}, Stdin: stdin, TerminalOutput: true}) + }() + select { + case err := <-done: + if !errors.Is(err, io.ErrNoProgress) || !strings.Contains(err.Error(), "stream output") { + t.Fatalf("got %v", err) + } + case <-time.After(time.Second): + t.Fatal("Run blocked after terminal output flush error") + } +} + +func TestReceiveDatagramsRejectsShortOutputWrite(t *testing.T) { + s := &fakeSession{received: [][]byte{{0, 1}}} + err := receiveDatagrams(context.Background(), Config{Session: s, Stdout: &shortWriter{}}) + if !errors.Is(err, io.ErrShortWrite) || !strings.Contains(err.Error(), "datagram output") { + t.Fatalf("got %v", err) + } +} + +func TestSendInputDatagramsHonorsCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + s := &fakeSession{} + err := sendInputDatagrams(ctx, s, strings.NewReader("payload"), core.WTDatagramBinary) + if !errors.Is(err, context.Canceled) { + t.Fatalf("got %v", err) + } + if len(s.datagrams) != 0 { + t.Fatalf("sent datagrams after cancellation: %#v", s.datagrams) + } +} + +func TestRunDatagramsDoesNotSendInitialDataAfterCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + for _, cfg := range []Config{ + {InitialPayloadSet: true, InitialPayload: []byte("payload")}, + {InitialReader: strings.NewReader("payload")}, + } { + s := &fakeSession{} + cfg.Session = s + cfg.Stdout = io.Discard + cfg.Mode = core.WTDatagram + err := Run(ctx, cfg) + if !errors.Is(err, context.Canceled) { + t.Fatalf("got %v", err) + } + if len(s.datagrams) != 0 { + t.Fatalf("sent datagrams after cancellation: %#v", s.datagrams) + } + } +} + +type peerCloseSession struct { + fakeSession + closed chan struct{} + once sync.Once + err error +} + +func (s *peerCloseSession) ReceiveDatagram(context.Context) ([]byte, error) { + return nil, s.err +} + +func (s *peerCloseSession) Close() error { + s.once.Do(func() { close(s.closed) }) + return nil +} + +type waitReader struct{ ready <-chan struct{} } + +func (r waitReader) Read(p []byte) (int, error) { + <-r.ready + return copy(p, "payload"), io.EOF +} + +func TestRunDatagramsPreservesPeerCloseError(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + peerErr := errors.New("peer closed session") + s := &peerCloseSession{fakeSession: fakeSession{}, closed: make(chan struct{}), err: peerErr} + err := Run(ctx, Config{Session: s, Stdout: io.Discard, Mode: core.WTDatagram, DatagramMode: core.WTDatagramBinary, Stdin: waitReader{ready: s.closed}}) + if !errors.Is(err, peerErr) { + t.Fatalf("got %v, want %v", err, peerErr) + } + if len(s.datagrams) != 0 { + t.Fatalf("sent datagrams after peer close: %#v", s.datagrams) + } +} + +func TestSendDatagramReportsCurrentQUICLimit(t *testing.T) { + s := &fakeSession{sendErr: &quic.DatagramTooLargeError{MaxDatagramPayloadSize: 1180}} + err := sendDatagram(s, make([]byte, 1200)) + if err == nil || !strings.Contains(err.Error(), "1200 bytes; current QUIC limit 1180 bytes before HTTP/3 overhead") { + t.Fatalf("got %v", err) + } +} + func TestTerminalWriterEscapesControlsAndKeepsUTF8(t *testing.T) { var out bytes.Buffer w := &terminalWriter{dst: &out} diff --git a/skills/fetch/references/webtransport.md b/skills/fetch/references/webtransport.md index 575cdf4e..15fe7f9d 100644 --- a/skills/fetch/references/webtransport.md +++ b/skills/fetch/references/webtransport.md @@ -10,6 +10,7 @@ raw when redirected and terminal-safe when displayed. `--wt-mode datagram` sends datagrams. `--wt-datagram-mode lines` sends one line per datagram, and `binary` sends 1 KiB chunks. Received datagrams are compact JSON Lines records containing `sequence`, `length`, and base64 `data`. +Oversized outgoing datagram errors include the current underlying QUIC limit. Input EOF does not close a datagram session; cancellation or peer closure does. Repeat `--wt-protocol` for application protocol offers. WebTransport v1 does