// w3cs WebTransport bridge: terminates browser QUIC sessions and bridges
// them onto a seat relay's local WebSocket planes.
//
// Why a sidecar: the relay already serves every command plane over local
// WebSockets (its TCP fallback), and loopback TCP never drops or reorders.
// The lossy leg is the internet path, and that is exactly the part QUIC
// replaces. The bridge re-planes the byte-identical relay messages:
//
//	frame plane     one QUIC unidirectional stream PER FRAME. A dependent
//	                frame superseded while its stream is still blocked is
//	                cancelled (RESET_STREAM), so one loss or stall can
//	                never head-of-line block later frames — the SCTP/TCP
//	                failure mode this bridge exists to remove. Keyframes
//	                and geometry anchors are never cancelled.
//	resource plane  one long-lived stream (ordered + lossless; texture XOR
//	                deltas are stateful and need exactly this).
//	recovery plane  one long-lived stream.
//	control (bidi)  hello/ready handshake, then reliable input events.
//	datagrams       pointer moves from the browser (latest-wins; a lost
//	                move is superseded by the next one ~16 ms later).
//
// Stream framing to the browser: first byte tags the plane (1 frame,
// 2 resource, 3 recovery), then length-prefixed relay messages exactly as
// the DataChannel would have delivered them. The proven client transport
// runs unchanged on top.
//
// Backpressure: each relay message is acked (text byte count, the same ack
// protocol the browser's WebSocket fallback speaks) only after its QUIC
// stream write returns, so relay-side in_flight reflects what QUIC flow
// control has not yet accepted and the relay governor keeps working.
package main

import (
	"context"
	"crypto/tls"
	"encoding/binary"
	"encoding/json"
	"flag"
	"fmt"
	"log"
	"net"
	"net/http"
	"net/url"
	"os"
	"strconv"
	"strings"
	"sync"
	"time"

	"github.com/coder/websocket"
	"github.com/quic-go/quic-go"
	"github.com/quic-go/quic-go/http3"
	"github.com/quic-go/webtransport-go"
)

const (
	tagFrame    = 1
	tagResource = 2
	tagRecovery = 3

	envelopeMagic   = 0x53433357 // "W3CS"
	envelopeVersion = 1
	kindFrame       = 3
	flagKeyframe    = 0x2
	flagAnchor      = 0x8

	// Cancel code for superseded dependent frames; the browser treats any
	// reset as "this frame never completed" and the Reassembler ages the
	// partial fragments out.
	cancelSuperseded = webtransport.StreamErrorCode(1)
)

var (
	listenAddr = flag.String("listen", ":4443", "HTTP/3 listen address")
	certFile   = flag.String("cert", "/etc/ssl/w3cs/w3cs.crt", "TLS certificate")
	keyFile    = flag.String("key", "/etc/ssl/w3cs/w3cs.key", "TLS key")
	relayHost  = flag.String("relay-host", "127.0.0.1", "relay host")
	relayBase  = flag.Int("relay-base", 8144, "relay port base (seat N = base+N)")
	seatCount  = flag.Int("seats", 9, "highest seat number served")
)

// certLoader re-reads the certificate when the file changes, so Let's
// Encrypt renewals apply without a restart.
type certLoader struct {
	mu       sync.Mutex
	cert     *tls.Certificate
	loadedAt time.Time
}

func (c *certLoader) get() (*tls.Certificate, error) {
	c.mu.Lock()
	defer c.mu.Unlock()
	if c.cert != nil && time.Since(c.loadedAt) < time.Minute {
		return c.cert, nil
	}
	cert, err := tls.LoadX509KeyPair(*certFile, *keyFile)
	if err != nil {
		if c.cert != nil {
			return c.cert, nil
		}
		return nil, err
	}
	c.cert = &cert
	c.loadedAt = time.Now()
	return c.cert, nil
}

type frameMsg struct {
	data     []byte
	frame    uint32
	reliable bool
	last     bool
	newFrame bool
}

// parseFrameMsg classifies one relay frame-plane message. Anything that is
// not a well-formed disposable frame fragment is treated as reliable: the
// bridge must never cancel what it does not understand.
func parseFrameMsg(data []byte) frameMsg {
	msg := frameMsg{data: data, reliable: true, newFrame: true, last: true}
	if len(data) < 32 ||
		binary.LittleEndian.Uint32(data) != envelopeMagic ||
		data[4] != envelopeVersion || data[5] != kindFrame {
		return msg
	}
	flags := binary.LittleEndian.Uint16(data[6:])
	msg.frame = binary.LittleEndian.Uint32(data[16:])
	fragIdx := binary.LittleEndian.Uint16(data[20:])
	fragCount := binary.LittleEndian.Uint16(data[22:])
	msg.reliable = flags&(flagKeyframe|flagAnchor) != 0
	msg.newFrame = fragIdx == 0
	msg.last = fragIdx == fragCount-1
	return msg
}

// How long a superseded dependent frame may keep writing before its
// stream is reset. Cancelling instantly on supersede looked right on
// paper but measured terribly: a transmitted-then-cancelled frame is
// pure wire waste, and under congestion (small cwnd, frame time close
// to the frame interval) it became a cancel spiral — the relay sent
// 27-29 fps while 3-17 fps completed arrival. Transmitted bytes are
// sunk cost: let the in-flight frame finish and apply latest-wins to
// the UNSTARTED queue instead (dropping those wastes nothing). The
// reset stays only as an emergency valve for a frame that is truly
// wedged while newer data waits.
const staleWriteCancel = 300 * time.Millisecond

// frameWriter owns the per-frame streams and the supersede policy. One
// goroutine writes; the pump goroutine enqueues, drops unstarted stale
// frames, and only resets a blocked stream past staleWriteCancel.
type frameWriter struct {
	session *webtransport.Session
	ack     func(int)

	mu         sync.Mutex
	cond       *sync.Cond
	queue      []frameMsg
	closed     bool
	cur        *webtransport.SendStream
	curFrame   uint32
	curRel     bool
	writing    bool
	writeStart time.Time
	aborted    map[uint32]bool

	// Session-lifetime counters plus a per-report peak, logged by the
	// session's stats reporter.
	opened       uint64
	finished     uint64
	cancelled    uint64
	queueDropped uint64
	written      uint64
	maxBlock     time.Duration
}

func newFrameWriter(session *webtransport.Session, ack func(int)) *frameWriter {
	writer := &frameWriter{session: session, ack: ack,
		aborted: make(map[uint32]bool)}
	writer.cond = sync.NewCond(&writer.mu)
	go writer.run()
	go writer.staleWatch()
	return writer
}

// maybeCancelStaleLocked resets the in-flight dependent frame only when a
// newer frame is waiting AND the write has been stuck past the stale
// window. Requires w.mu.
func (w *frameWriter) maybeCancelStaleLocked(newerFrame uint32) {
	if w.writing && !w.curRel && w.cur != nil && w.curFrame < newerFrame &&
		time.Since(w.writeStart) > staleWriteCancel {
		w.aborted[w.curFrame] = true
		w.cancelled++
		w.cur.CancelWrite(cancelSuperseded)
	}
}

// staleWatch backs up the enqueue-time check: once the relay pauses under
// backpressure, no new enqueue would ever fire the stale cancel for the
// frame that is blocking everything queued behind it.
func (w *frameWriter) staleWatch() {
	ticker := time.NewTicker(100 * time.Millisecond)
	defer ticker.Stop()
	for range ticker.C {
		w.mu.Lock()
		if w.closed {
			w.mu.Unlock()
			return
		}
		newest := uint32(0)
		for _, queued := range w.queue {
			if queued.frame > newest {
				newest = queued.frame
			}
		}
		if newest > 0 {
			w.maybeCancelStaleLocked(newest)
		}
		w.mu.Unlock()
	}
}

func (w *frameWriter) enqueue(msg frameMsg) {
	w.mu.Lock()
	defer w.mu.Unlock()
	if w.closed {
		w.ack(len(msg.data))
		return
	}
	if msg.newFrame && !msg.reliable {
		// Latest-wins on the unstarted queue: dropping a frame that never
		// started costs no wire. Fragments of the frame currently being
		// written stay — dropping them would leave its stream dangling.
		kept := w.queue[:0]
		for _, queued := range w.queue {
			if !queued.reliable && queued.frame < msg.frame &&
				queued.frame != w.curFrame {
				w.aborted[queued.frame] = true
				w.queueDropped++
				w.ack(len(queued.data))
				continue
			}
			kept = append(kept, queued)
		}
		w.queue = kept
		w.maybeCancelStaleLocked(msg.frame)
	}
	w.queue = append(w.queue, msg)
	w.cond.Signal()
}

func (w *frameWriter) close() {
	w.mu.Lock()
	w.closed = true
	if w.cur != nil {
		w.cur.CancelWrite(cancelSuperseded)
	}
	w.cond.Broadcast()
	w.mu.Unlock()
}

func (w *frameWriter) run() {
	for {
		w.mu.Lock()
		for len(w.queue) == 0 && !w.closed {
			w.cond.Wait()
		}
		if w.closed {
			for _, queued := range w.queue {
				w.ack(len(queued.data))
			}
			w.queue = nil
			w.mu.Unlock()
			return
		}
		msg := w.queue[0]
		w.queue = w.queue[1:]
		if w.aborted[msg.frame] && !msg.reliable {
			if msg.last {
				delete(w.aborted, msg.frame)
			}
			w.ack(len(msg.data))
			w.mu.Unlock()
			continue
		}
		needStream := w.cur == nil || w.curFrame != msg.frame
		if needStream && w.cur != nil {
			// The relay moved on without a last fragment (it can drop a
			// frame mid-send under its own backpressure); finish the old
			// stream so the browser sees a clean end, not a stall.
			w.cur.Close()
			w.cur = nil
		}
		w.mu.Unlock()

		if needStream {
			stream, err := w.session.OpenUniStreamSync(w.session.Context())
			if err != nil {
				w.ack(len(msg.data))
				continue
			}
			if _, err := stream.Write([]byte{tagFrame}); err != nil {
				stream.CancelWrite(cancelSuperseded)
				w.ack(len(msg.data))
				continue
			}
			w.mu.Lock()
			w.cur = stream
			w.curFrame = msg.frame
			w.curRel = msg.reliable
			w.writeStart = time.Now()
			w.opened++
			w.mu.Unlock()
		}

		w.mu.Lock()
		stream := w.cur
		w.writing = true
		w.mu.Unlock()
		writeBegan := time.Now()
		ok := stream != nil
		if ok {
			buffer := make([]byte, 4+len(msg.data))
			binary.LittleEndian.PutUint32(buffer, uint32(len(msg.data)))
			copy(buffer[4:], msg.data)
			if _, err := stream.Write(buffer); err != nil {
				ok = false
			}
		}
		w.mu.Lock()
		w.writing = false
		if blocked := time.Since(writeBegan); blocked > w.maxBlock {
			w.maxBlock = blocked
		}
		if !ok {
			// Cancelled (stale) or session-dead: drop the rest of this
			// frame as it pops.
			w.aborted[msg.frame] = true
			if w.cur == stream {
				w.cur = nil
			}
		} else {
			w.written++
			if msg.last {
				stream.Close()
				if w.cur == stream {
					w.cur = nil
				}
				w.finished++
				delete(w.aborted, msg.frame)
			}
		}
		// The aborted set only needs to cover frames still draining out of
		// the queue; keep it from growing over a long session.
		if len(w.aborted) > 64 {
			w.aborted = make(map[uint32]bool)
		}
		w.mu.Unlock()
		w.ack(len(msg.data))
	}
}

// snapshotStats returns the lifetime counters plus the max single-write
// block since the previous snapshot.
func (w *frameWriter) snapshotStats() (opened, finished, cancelled,
	queueDropped uint64, queueLen int, maxBlock time.Duration) {
	w.mu.Lock()
	defer w.mu.Unlock()
	opened, finished = w.opened, w.finished
	cancelled, queueDropped = w.cancelled, w.queueDropped
	queueLen = len(w.queue)
	maxBlock = w.maxBlock
	w.maxBlock = 0
	return
}

// relayPlane is one WebSocket to the relay plus its serialized ack sender.
type relayPlane struct {
	conn  *websocket.Conn
	ackCh chan int
}

func dialPlane(ctx context.Context, base, path, session string) (*relayPlane, error) {
	dialCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
	defer cancel()
	conn, _, err := websocket.Dial(dialCtx,
		fmt.Sprintf("%s/%s?session=%s", base, path, session), nil)
	if err != nil {
		return nil, err
	}
	conn.SetReadLimit(2 * 1024 * 1024)
	plane := &relayPlane{conn: conn, ackCh: make(chan int, 1024)}
	go func() {
		for bytes := range plane.ackCh {
			writeCtx, done := context.WithTimeout(ctx, 10*time.Second)
			err := conn.Write(writeCtx, websocket.MessageText,
				[]byte(strconv.Itoa(bytes)))
			done()
			if err != nil {
				return
			}
		}
	}()
	return plane, nil
}

func (p *relayPlane) ack(bytes int) {
	select {
	case p.ackCh <- bytes:
	default:
		// The ack channel only backs up if the relay stopped reading; the
		// in-flight counter recovers at the next plane connect.
	}
}

type seatSession struct {
	seat    int
	session *webtransport.Session
	cancel  context.CancelFunc
}

var (
	activeMu sync.Mutex
	active   = make(map[int]*seatSession)
)

func lengthPrefixed(data []byte) []byte {
	out := make([]byte, 4+len(data))
	binary.LittleEndian.PutUint32(out, uint32(len(data)))
	copy(out[4:], data)
	return out
}

// pumpReliable copies one relay plane onto one long-lived uni stream.
func pumpReliable(ctx context.Context, plane *relayPlane,
	session *webtransport.Session, tag byte, fail func(error)) {
	stream, err := session.OpenUniStreamSync(ctx)
	if err != nil {
		fail(fmt.Errorf("open stream %d: %w", tag, err))
		return
	}
	if _, err := stream.Write([]byte{tag}); err != nil {
		fail(fmt.Errorf("stream %d tag: %w", tag, err))
		return
	}
	for {
		kind, data, err := plane.conn.Read(ctx)
		if err != nil {
			fail(fmt.Errorf("relay plane %d read: %w", tag, err))
			return
		}
		if kind != websocket.MessageBinary {
			continue
		}
		if _, err := stream.Write(lengthPrefixed(data)); err != nil {
			fail(fmt.Errorf("stream %d write: %w", tag, err))
			return
		}
		plane.ack(len(data))
	}
}

func pumpFrames(ctx context.Context, plane *relayPlane,
	writer *frameWriter, fail func(error)) {
	for {
		kind, data, err := plane.conn.Read(ctx)
		if err != nil {
			fail(fmt.Errorf("relay frame read: %w", err))
			return
		}
		if kind != websocket.MessageBinary {
			continue
		}
		msg := parseFrameMsg(data)
		writer.enqueue(msg)
	}
}

func serveSession(seat int, session *webtransport.Session, sessionID string) {
	ctx, cancel := context.WithCancel(session.Context())
	defer cancel()

	// One session per seat, newest wins — the relay applies the same rule
	// to its planes, so a stale bridge must not hold them.
	activeMu.Lock()
	if old := active[seat]; old != nil {
		old.session.CloseWithError(0, "seat session replaced")
		old.cancel()
	}
	current := &seatSession{seat: seat, session: session, cancel: cancel}
	active[seat] = current
	activeMu.Unlock()
	defer func() {
		activeMu.Lock()
		if active[seat] == current {
			delete(active, seat)
		}
		activeMu.Unlock()
	}()

	// The control stream opens first: its hello triggers the plane dials
	// and its ready reply tells the page it can start the game session.
	control, err := session.AcceptStream(ctx)
	if err != nil {
		return
	}
	controlWriteMu := sync.Mutex{}
	sendControl := func(object any) {
		payload, err := json.Marshal(object)
		if err != nil {
			return
		}
		controlWriteMu.Lock()
		defer controlWriteMu.Unlock()
		control.Write(lengthPrefixed(payload))
	}

	readControl := func() ([]byte, error) {
		header := make([]byte, 4)
		if _, err := ioReadFull(control, header); err != nil {
			return nil, err
		}
		size := binary.LittleEndian.Uint32(header)
		if size == 0 || size > 1024*1024 {
			return nil, fmt.Errorf("invalid control length %d", size)
		}
		payload := make([]byte, size)
		if _, err := ioReadFull(control, payload); err != nil {
			return nil, err
		}
		return payload, nil
	}

	hello, err := readControl()
	if err != nil {
		return
	}
	var helloMsg struct {
		T       string `json:"t"`
		Session string `json:"session"`
	}
	if json.Unmarshal(hello, &helloMsg) != nil || helloMsg.T != "hello" {
		sendControl(map[string]any{"t": "wtError", "message": "expected hello"})
		return
	}
	if sessionID == "" {
		sessionID = helloMsg.Session
	}

	base := fmt.Sprintf("ws://%s:%d", *relayHost, *relayBase+seat)
	resource, err := dialPlane(ctx, base, "resource", sessionID)
	if err != nil {
		sendControl(map[string]any{"t": "wtError",
			"message": "relay resource plane unavailable"})
		return
	}
	defer resource.conn.CloseNow()
	recovery, err := dialPlane(ctx, base, "recovery", sessionID)
	if err != nil {
		sendControl(map[string]any{"t": "wtError",
			"message": "relay recovery plane unavailable"})
		return
	}
	defer recovery.conn.CloseNow()
	frame, err := dialPlane(ctx, base, "frame", sessionID)
	if err != nil {
		sendControl(map[string]any{"t": "wtError",
			"message": "relay frame plane unavailable"})
		return
	}
	defer frame.conn.CloseNow()
	// The input lane is additive (the page keeps its signaling copy), so an
	// older relay without /input degrades gracefully.
	input, inputErr := dialPlane(ctx, base, "input", sessionID)
	if input != nil {
		defer input.conn.CloseNow()
	}

	log.Printf("seat %d session %s: planes attached (input=%v)",
		seat, sessionID, inputErr == nil)
	sendControl(map[string]any{"t": "wtReady", "seat": seat,
		"input": inputErr == nil})

	var failOnce sync.Once
	fail := func(err error) {
		failOnce.Do(func() {
			log.Printf("seat %d session %s: %v", seat, sessionID, err)
			session.CloseWithError(0, "bridge plane failed")
			cancel()
		})
	}

	writer := newFrameWriter(session, frame.ack)
	defer writer.close()

	// Frame-plane health, the counters that caught the cancel spiral:
	// finished must track opened, cancels/queue-drops must stay rare, and
	// maxBlockMs is how hard QUIC flow control pushed back this interval.
	go func() {
		ticker := time.NewTicker(5 * time.Second)
		defer ticker.Stop()
		for {
			select {
			case <-ctx.Done():
				return
			case <-ticker.C:
				opened, finished, cancelled, queueDropped, queueLen,
					maxBlock := writer.snapshotStats()
				log.Printf("seat %d wt-frames opened=%d finished=%d "+
					"cancelled=%d qdrop=%d queue=%d maxBlockMs=%d",
					seat, opened, finished, cancelled, queueDropped,
					queueLen, maxBlock.Milliseconds())
			}
		}
	}()

	go pumpReliable(ctx, resource, session, tagResource, fail)
	go pumpReliable(ctx, recovery, session, tagRecovery, fail)
	go pumpFrames(ctx, frame, writer, fail)

	forwardInput := func(event []byte) {
		if input == nil {
			return
		}
		writeCtx, done := context.WithTimeout(ctx, 5*time.Second)
		defer done()
		if err := input.conn.Write(writeCtx, websocket.MessageText,
			event); err != nil {
			input = nil
		}
	}

	// Datagrams: pointer moves. Loss needs no handling — the page races
	// every input over its reliable lanes too and the relay dedupes.
	go func() {
		for {
			event, err := session.ReceiveDatagram(ctx)
			if err != nil {
				return
			}
			forwardInput(event)
		}
	}()

	// Drain the input plane so relay-side pings never back up (the relay
	// does not send on it today).
	if input != nil {
		go func(plane *relayPlane) {
			for {
				if _, _, err := plane.conn.Read(ctx); err != nil {
					return
				}
			}
		}(input)
	}

	// Control stream: reliable input events until the session ends.
	for {
		event, err := readControl()
		if err != nil {
			return
		}
		forwardInput(event)
	}
}

// ioReadFull avoids importing io just for ReadFull semantics on the
// webtransport stream interface.
func ioReadFull(reader interface{ Read([]byte) (int, error) },
	buffer []byte) (int, error) {
	total := 0
	for total < len(buffer) {
		n, err := reader.Read(buffer[total:])
		total += n
		if err != nil {
			return total, err
		}
	}
	return total, nil
}

func main() {
	flag.Parse()
	log.SetFlags(log.LstdFlags | log.LUTC)

	loader := &certLoader{}
	if _, err := loader.get(); err != nil {
		log.Fatalf("load certificate: %v", err)
	}
	tlsConf := &tls.Config{
		// quic-go does not add the h3 ALPN to a custom TLSConfig; without
		// it every handshake dies with "no ALPN protocol selected".
		NextProtos: []string{http3.NextProtoH3},
		GetCertificate: func(*tls.ClientHelloInfo) (*tls.Certificate, error) {
			return loader.get()
		},
	}

	wtServer := &webtransport.Server{
		// The page lives on :443 and this bridge on :4443, so the default
		// same-authority origin check always fails. Same hostname is the
		// real trust boundary here; localhost stays open for lab pages.
		CheckOrigin: func(r *http.Request) bool {
			origin := r.Header.Get("Origin")
			if origin == "" {
				return true
			}
			parsed, err := url.Parse(origin)
			if err != nil {
				return false
			}
			originHost := parsed.Hostname()
			requestHost := r.Host
			if host, _, err := net.SplitHostPort(requestHost); err == nil {
				requestHost = host
			}
			return strings.EqualFold(originHost, requestHost) ||
				originHost == "localhost" || originHost == "127.0.0.1"
		},
		H3: &http3.Server{
			Addr:            *listenAddr,
			TLSConfig:       tlsConf,
			EnableDatagrams: true,
			QUICConfig: &quic.Config{
				EnableDatagrams: true,
				MaxIdleTimeout:  45 * time.Second,
			},
		},
	}

	mux := http.NewServeMux()
	mux.HandleFunc("/seat/", func(w http.ResponseWriter, r *http.Request) {
		rest := strings.TrimPrefix(r.URL.Path, "/seat/")
		seat, err := strconv.Atoi(rest)
		if err != nil || seat < 1 || seat > *seatCount {
			http.Error(w, "unknown seat", http.StatusNotFound)
			return
		}
		sessionID := r.URL.Query().Get("session")
		session, err := wtServer.Upgrade(w, r)
		if err != nil {
			log.Printf("seat %d upgrade failed: %v", seat, err)
			http.Error(w, "upgrade failed", http.StatusBadRequest)
			return
		}
		log.Printf("seat %d session %s connected from %s (origin %q)",
			seat, sessionID, r.RemoteAddr, r.Header.Get("Origin"))
		go serveSession(seat, session, sessionID)
	})
	wtServer.H3.Handler = mux

	log.Printf("w3cs wt-bridge listening on %s (relay %s:%d+N)",
		*listenAddr, *relayHost, *relayBase)
	if err := wtServer.ListenAndServe(); err != nil {
		log.Fatalf("listen: %v", err)
		os.Exit(1)
	}
}
