webtransport API

webtransport

package

API reference for the webtransport package.

S
struct

Handler

Handler upgrades HTTP/3 requests and binds WebTransport sessions to a Bus.

streambus/webtransport/handler.go:18-27
type Handler struct

Methods

ServeHTTP
Method

Parameters

func (*Handler) ServeHTTP(writer http.ResponseWriter, request *http.Request)
{
	if h.Server == nil || h.Bus == nil {
		http.Error(writer, "streambus webtransport is not configured", http.StatusServiceUnavailable)
		return
	}
	if h.Authenticate != nil {
		if err := h.Authenticate(request); err != nil {
			http.Error(writer, "unauthorized", http.StatusUnauthorized)
			return
		}
	}
	session, err := h.Server.Upgrade(writer, request)
	if err != nil {
		h.report(err)
		return
	}
	if err := h.ServeSession(session.Context(), session); err != nil && !errors.Is(err, context.Canceled) {
		h.report(err)
	}
}
ServeSession
Method

ServeSession runs the StreamBus protocol over an upgraded session.

Parameters

session *wt.Session

Returns

error
func (*Handler) ServeSession(parent context.Context, session *wt.Session) error
{
	if h.Bus == nil {
		return errors.New("streambus webtransport: nil bus")
	}
	ctx, cancel := context.WithCancel(parent)
	defer cancel()
	defer session.CloseWithError(0, "")
	control, err := session.AcceptStream(ctx)
	if err != nil {
		return fmt.Errorf("accept control stream: %w", err)
	}
	defer control.Close()

	state := &sessionState{
		handler:       h,
		session:       session,
		control:       control,
		subscriptions: make(map[uint64]*streambus.Subscription),
		acks:          make(map[uint64]uint64),
	}
	defer state.closeSubscriptions()

	errCh := make(chan error, 2)
	go func() { errCh <- state.readControl(ctx) }()
	workers := 1
	datagrams := session.SessionState().ConnectionState.SupportsDatagrams
	if datagrams.Local && datagrams.Remote {
		workers++
		go func() { errCh <- state.readDatagrams(ctx) }()
	}
	err = <-errCh
	cancel()
	_ = session.CloseWithError(0, "")
	for i := 1; i < workers; i++ {
		<-errCh
	}
	if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
		return nil
	}
	return err
}
maximum
Method

Returns

int
func (*Handler) maximum() int
{
	if h.MaxMessageBytes <= 0 {
		return defaultMaxMessageBytes
	}
	return h.MaxMessageBytes
}

Returns

int
func (*Handler) maximumDatagram() int
{
	if h.MaxDatagramBytes <= 0 {
		return defaultMaxDatagramBytes
	}
	return h.MaxDatagramBytes
}
report
Method

Parameters

err error
func (*Handler) report(err error)
{
	if h.OnError != nil {
		h.OnError(err)
	}
}

Fields

Name Type Description
Server *wt.Server
Bus streambus.Bus
Authenticate func(*http.Request) error
MaxMessageBytes int
MaxDatagramBytes int
OnError func(error)
S
struct

sessionState

streambus/webtransport/handler.go:113-121
type sessionState struct

Methods

readControl
Method

Parameters

Returns

error
func (*sessionState) readControl(ctx context.Context) error
{
	reader := streambus.NewReader(s.control, s.handler.maximum())
	for {
		message, err := reader.ReadMessage()
		if err != nil {
			return err
		}
		if err := s.handleMessage(ctx, message, false); err != nil {
			if writeErr := s.writeControl(streambus.Message{Kind: streambus.MessageError, ID: message.ID, Error: err.Error()}); writeErr != nil {
				return writeErr
			}
		}
	}
}
readDatagrams
Method

Parameters

Returns

error
func (*sessionState) readDatagrams(ctx context.Context) error
{
	for {
		data, err := s.session.ReceiveDatagram(ctx)
		if err != nil {
			return err
		}
		if len(data) > s.handler.maximum() {
			continue
		}
		message, err := streambus.DecodeMessage(data)
		if err != nil {
			continue
		}
		if err := s.handleMessage(ctx, message, true); err != nil {
			s.handler.report(err)
		}
	}
}
handleMessage
Method

Parameters

datagram bool

Returns

error
func (*sessionState) handleMessage(ctx context.Context, message streambus.Message, datagram bool) error
{
	switch message.Kind {
	case streambus.MessageSubscribe:
		if datagram {
			return errors.New("subscribe requires the reliable control stream")
		}
		return s.subscribe(ctx, message)
	case streambus.MessageUnsubscribe:
		if datagram {
			return errors.New("unsubscribe requires the reliable control stream")
		}
		return s.unsubscribe(message.SubscriptionID)
	case streambus.MessagePublish:
		if datagram {
			message.Frame.Reliability = streambus.Unreliable
		}
		sequence, err := s.handler.Bus.Publish(ctx, message.Frame)
		if err != nil {
			return err
		}
		if !datagram {
			return s.writeControl(streambus.Message{Kind: streambus.MessageAck, ID: message.ID, Ack: sequence})
		}
		return nil
	case streambus.MessageAck:
		if datagram {
			return errors.New("ack requires the reliable control stream")
		}
		s.subsMu.Lock()
		if _, exists := s.subscriptions[message.SubscriptionID]; exists {
			s.acks[message.SubscriptionID] = message.Ack
		}
		s.subsMu.Unlock()
		return nil
	case streambus.MessagePing:
		if datagram {
			data, err := streambus.EncodeMessage(streambus.Message{Kind: streambus.MessagePong, ID: message.ID})
			if err != nil {
				return err
			}
			return s.session.SendDatagram(data)
		}
		return s.writeControl(streambus.Message{Kind: streambus.MessagePong, ID: message.ID})
	default:
		return fmt.Errorf("unsupported message kind %d", message.Kind)
	}
}
subscribe
Method

Parameters

Returns

error
func (*sessionState) subscribe(ctx context.Context, message streambus.Message) error
{
	if message.SubscriptionID == 0 {
		return errors.New("subscription id is required")
	}
	s.subsMu.Lock()
	if _, exists := s.subscriptions[message.SubscriptionID]; exists {
		s.subsMu.Unlock()
		return errors.New("subscription id already exists")
	}
	s.subsMu.Unlock()

	options := message.Options
	if options.Topic == "" {
		options.Topic = message.Frame.Topic
	}
	subscription, err := s.handler.Bus.Subscribe(ctx, options)
	if err != nil {
		return err
	}
	s.subsMu.Lock()
	s.subscriptions[message.SubscriptionID] = subscription
	s.acks[message.SubscriptionID] = message.Options.Since
	s.subsMu.Unlock()
	go func() {
		if err := s.pump(ctx, message.SubscriptionID, subscription); err != nil && !errors.Is(err, context.Canceled) {
			s.handler.report(err)
		}
	}()
	return s.writeControl(streambus.Message{Kind: streambus.MessageAck, ID: message.ID, SubscriptionID: message.SubscriptionID})
}
unsubscribe
Method

Parameters

id uint64

Returns

error
func (*sessionState) unsubscribe(id uint64) error
{
	s.subsMu.Lock()
	subscription := s.subscriptions[id]
	delete(s.subscriptions, id)
	delete(s.acks, id)
	s.subsMu.Unlock()
	if subscription == nil {
		return errors.New("unknown subscription")
	}
	return subscription.Close()
}
pump
Method

Parameters

id uint64
subscription *streambus.Subscription

Returns

error
func (*sessionState) pump(ctx context.Context, id uint64, subscription *streambus.Subscription) error
{
	defer func() {
		s.subsMu.Lock()
		delete(s.subscriptions, id)
		delete(s.acks, id)
		s.subsMu.Unlock()
		_ = subscription.Close()
	}()

	var reliable *wt.SendStream
	for {
		select {
		case <-ctx.Done():
			if reliable != nil {
				_ = reliable.Close()
			}
			return ctx.Err()
		case frame, ok := <-subscription.Frames():
			if !ok {
				if reliable != nil {
					_ = reliable.Close()
				}
				return nil
			}
			message := streambus.Message{Kind: streambus.MessageData, SubscriptionID: id, Frame: frame}
			if frame.Reliability == streambus.Unreliable {
				data, err := streambus.EncodeMessage(message)
				if err != nil {
					return err
				}
				state := s.session.SessionState().ConnectionState.SupportsDatagrams
				if len(data) <= s.handler.maximumDatagram() && state.Local && state.Remote {
					if err := s.session.SendDatagram(data); err == nil {
						continue
					}
				}
			}
			if reliable == nil {
				stream, err := s.session.OpenUniStreamSync(ctx)
				if err != nil {
					return err
				}
				reliable = stream
				if err := streambus.WriteMessage(reliable, streambus.Message{Kind: streambus.MessageStreamOpen, SubscriptionID: id}); err != nil {
					_ = reliable.Close()
					return err
				}
			}
			if err := streambus.WriteMessage(reliable, message); err != nil {
				_ = reliable.Close()
				return err
			}
		}
	}
}
writeControl
Method

Parameters

Returns

error
func (*sessionState) writeControl(message streambus.Message) error
{
	s.writeMu.Lock()
	defer s.writeMu.Unlock()
	return streambus.WriteMessage(s.control, message)
}
func (*sessionState) closeSubscriptions()
{
	s.subsMu.Lock()
	subscriptions := make([]*streambus.Subscription, 0, len(s.subscriptions))
	for _, subscription := range s.subscriptions {
		subscriptions = append(subscriptions, subscription)
	}
	s.subscriptions = make(map[uint64]*streambus.Subscription)
	s.acks = make(map[uint64]uint64)
	s.subsMu.Unlock()
	for _, subscription := range subscriptions {
		_ = subscription.Close()
	}
}

Fields

Name Type Description
handler *Handler
session *wt.Session
control *wt.Stream
writeMu sync.Mutex
subsMu sync.Mutex
subscriptions map[uint64]*streambus.Subscription
acks map[uint64]uint64
F
function

TestHandlerRequiresConfiguration

Parameters

streambus/webtransport/handler_test.go:14-21
func TestHandlerRequiresConfiguration(t *testing.T)

{
	recorder := httptest.NewRecorder()
	request := httptest.NewRequest(http.MethodConnect, "https://example.test/stream", nil)
	(&Handler{}).ServeHTTP(recorder, request)
	if recorder.Code != http.StatusServiceUnavailable {
		t.Fatalf("status = %d", recorder.Code)
	}
}
F
function

TestHandlerAuthenticatesBeforeUpgrade

Parameters

streambus/webtransport/handler_test.go:23-39
func TestHandlerAuthenticatesBeforeUpgrade(t *testing.T)

{
	bus := streambus.NewInMemory(streambus.Config{})
	t.Cleanup(func() { _ = bus.Close() })
	handler := &Handler{
		Server: &wt.Server{},
		Bus:    bus,
		Authenticate: func(*http.Request) error {
			return errors.New("denied")
		},
	}
	recorder := httptest.NewRecorder()
	request := httptest.NewRequest(http.MethodConnect, "https://example.test/stream", nil)
	handler.ServeHTTP(recorder, request)
	if recorder.Code != http.StatusUnauthorized {
		t.Fatalf("status = %d", recorder.Code)
	}
}
F
function

TestDefaults

Parameters

streambus/webtransport/handler_test.go:41-49
func TestDefaults(t *testing.T)

{
	handler := &Handler{}
	if handler.maximum() != defaultMaxMessageBytes {
		t.Fatalf("maximum = %d", handler.maximum())
	}
	if handler.maximumDatagram() != defaultMaxDatagramBytes {
		t.Fatalf("maximum datagram = %d", handler.maximumDatagram())
	}
}
F
function

TestClientAckIsTracked

Parameters

streambus/webtransport/handler_test.go:51-70
func TestClientAckIsTracked(t *testing.T)

{
	bus := streambus.NewInMemory(streambus.Config{})
	t.Cleanup(func() { _ = bus.Close() })
	subscription, err := bus.Subscribe(context.Background(), streambus.SubscribeOptions{Topic: "x", Buffer: 1})
	if err != nil {
		t.Fatal(err)
	}
	state := &sessionState{
		subscriptions: map[uint64]*streambus.Subscription{7: subscription},
		acks:          make(map[uint64]uint64),
	}
	if err := state.handleMessage(context.Background(), streambus.Message{
		Kind: streambus.MessageAck, SubscriptionID: 7, Ack: 42,
	}, false); err != nil {
		t.Fatal(err)
	}
	if state.acks[7] != 42 {
		t.Fatalf("ack = %d", state.acks[7])
	}
}
F
function

TestWebTransportReliableAndDatagramDelivery

Parameters

streambus/webtransport/integration_test.go:23-145
func TestWebTransportReliableAndDatagramDelivery(t *testing.T)

{
	serverTLS := testTLSConfig(t)
	bus := streambus.NewInMemory(streambus.Config{ReplayCapacity: 8})
	t.Cleanup(func() { _ = bus.Close() })

	mux := http.NewServeMux()
	h3 := &http3.Server{
		TLSConfig: serverTLS,
		QUICConfig: &quic.Config{
			EnableDatagrams:                  true,
			EnableStreamResetPartialDelivery: true,
		},
		EnableDatagrams: true,
		Handler:         mux,
	}
	server := &wt.Server{H3: h3}
	wt.ConfigureHTTP3Server(h3)
	mux.Handle("/stream", &Handler{Server: server, Bus: bus})

	address, err := net.ResolveUDPAddr("udp", "127.0.0.1:0")
	if err != nil {
		t.Fatal(err)
	}
	packetConn, err := net.ListenUDP("udp", address)
	if err != nil {
		t.Fatal(err)
	}
	serveDone := make(chan error, 1)
	go func() { serveDone <- server.Serve(packetConn) }()
	t.Cleanup(func() {
		_ = server.Close()
		_ = packetConn.Close()
		select {
		case <-serveDone:
		case <-time.After(time.Second):
			t.Error("WebTransport server did not stop")
		}
	})

	dialer := &wt.Dialer{
		TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, // test-only certificate
		QUICConfig: &quic.Config{
			EnableDatagrams:                  true,
			EnableStreamResetPartialDelivery: true,
		},
	}
	t.Cleanup(func() { _ = dialer.Close() })
	ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
	defer cancel()
	url := fmt.Sprintf("https://localhost:%d/stream", packetConn.LocalAddr().(*net.UDPAddr).Port)
	response, session, err := dialer.Dial(ctx, url, nil)
	if err != nil {
		t.Fatal(err)
	}
	if response.StatusCode != http.StatusOK {
		t.Fatalf("status = %d", response.StatusCode)
	}
	t.Cleanup(func() { _ = session.CloseWithError(0, "") })

	control, err := session.OpenStreamSync(ctx)
	if err != nil {
		t.Fatal(err)
	}
	if err := streambus.WriteMessage(control, streambus.Message{
		Kind:           streambus.MessageSubscribe,
		ID:             1,
		SubscriptionID: 7,
		Options: streambus.SubscribeOptions{
			Topic: "ui:metrics", Buffer: 8, Overflow: streambus.LatestOnly,
		},
	}); err != nil {
		t.Fatal(err)
	}
	ack, err := streambus.NewReader(control, 4096).ReadMessage()
	if err != nil {
		t.Fatal(err)
	}
	if ack.Kind != streambus.MessageAck || ack.SubscriptionID != 7 {
		t.Fatalf("subscribe ack = %#v", ack)
	}

	if _, err := bus.Publish(ctx, streambus.Frame{
		Topic: "ui:metrics", Payload: []byte("snapshot"), Reliability: streambus.Reliable,
	}); err != nil {
		t.Fatal(err)
	}
	reliable, err := session.AcceptUniStream(ctx)
	if err != nil {
		t.Fatal(err)
	}
	reliableReader := streambus.NewReader(reliable, 4096)
	opened, err := reliableReader.ReadMessage()
	if err != nil {
		t.Fatal(err)
	}
	if opened.Kind != streambus.MessageStreamOpen || opened.SubscriptionID != 7 {
		t.Fatalf("stream open = %#v", opened)
	}
	data, err := reliableReader.ReadMessage()
	if err != nil {
		t.Fatal(err)
	}
	if data.Kind != streambus.MessageData || string(data.Frame.Payload) != "snapshot" {
		t.Fatalf("reliable data = %#v", data)
	}

	if _, err := bus.Publish(ctx, streambus.Frame{
		Topic: "ui:metrics", Payload: []byte("cursor"), Reliability: streambus.Unreliable,
	}); err != nil {
		t.Fatal(err)
	}
	datagram, err := session.ReceiveDatagram(ctx)
	if err != nil {
		t.Fatal(err)
	}
	message, err := streambus.DecodeMessage(datagram)
	if err != nil {
		t.Fatal(err)
	}
	if message.Kind != streambus.MessageData || string(message.Frame.Payload) != "cursor" {
		t.Fatalf("datagram data = %#v", message)
	}
}
F
function

testTLSConfig

Parameters

Returns

streambus/webtransport/integration_test.go:147-172
func testTLSConfig(t *testing.T) *tls.Config

{
	t.Helper()
	key, err := rsa.GenerateKey(rand.Reader, 2048)
	if err != nil {
		t.Fatal(err)
	}
	template := &x509.Certificate{
		SerialNumber: big.NewInt(1),
		Subject:      pkix.Name{CommonName: "localhost"},
		NotBefore:    time.Now().Add(-time.Minute),
		NotAfter:     time.Now().Add(time.Hour),
		KeyUsage:     x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
		ExtKeyUsage:  []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
		DNSNames:     []string{"localhost"},
		IPAddresses:  []net.IP{net.ParseIP("127.0.0.1")},
	}
	der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
	if err != nil {
		t.Fatal(err)
	}
	certificate := tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key}
	return &tls.Config{
		Certificates: []tls.Certificate{certificate},
		NextProtos:   []string{http3.NextProtoH3},
	}
}