package vk import ( "context" "encoding/binary" "fmt" "io" "net" "sync" "time" "github.com/pion/rtp" "github.com/pion/rtp/codecs" "github.com/pion/webrtc/v4" "github.com/sagernet/sing-box/transport/call/common" "github.com/sagernet/sing-box/transport/call/tunnel" "github.com/sagernet/sing/common/logger" M "github.com/sagernet/sing/common/metadata" N "github.com/sagernet/sing/common/network" ) type Relay interface { Init(iceServers []webrtc.ICEServer) error CreateOffer() (webrtc.SessionDescription, error) CreateAnswer() (webrtc.SessionDescription, error) SetRemoteDescription(sdpType webrtc.SDPType, sdp string) error AddICECandidate(candidate webrtc.ICECandidateInit) error OnICECandidate(fn func(*webrtc.ICECandidate)) OnConnectionStateChange(fn func(webrtc.PeerConnectionState)) Close() } type dcConn struct { conn net.Conn ch chan []byte } type TunnelRelay struct { pc *webrtc.PeerConnection remoteSet bool pending []webrtc.ICECandidateInit externalICE func(*webrtc.ICECandidate) externalCSC func(webrtc.PeerConnectionState) dc *webrtc.DataChannel dcMu sync.Mutex conns sync.Map sampleTrack *webrtc.TrackLocalStaticSample tun *tunnel.VP8DataTunnel obf *tunnel.TunnelObfuscator OnConnected func(tunnel.DataTunnel) screenDC *webrtc.DataChannel producerScreen *webrtc.DataChannel sym *tunnel.SymmetricScreenTunnel dialer N.Dialer readBufSize int logger logger.ContextLogger mode string modeOnce sync.Once } func NewTunnelRelay(dialer N.Dialer, logger logger.ContextLogger) *TunnelRelay { return &TunnelRelay{mode: "unknown", dialer: dialer, logger: logger} } func (u *TunnelRelay) SetObfuscator(o *tunnel.TunnelObfuscator) { u.obf = o } func (u *TunnelRelay) Init(iceServers []webrtc.ICEServer) error { pc, err := webrtc.NewPeerConnection(webrtc.Configuration{ICEServers: iceServers}) if err != nil { return err } u.pc = pc negotiated := true dcID := uint16(2) dc, err := pc.CreateDataChannel("tunnel", &webrtc.DataChannelInit{ Negotiated: &negotiated, ID: &dcID, }) if err != nil { u.logger.Warn(fmt.Sprintf("[relay] could not create tunnel DC: %v", err)) } else { u.dc = dc dc.OnOpen(func() { u.logger.Debug(fmt.Sprintf("[relay] tunnel DC open (readyState=%v)", dc.ReadyState())) }) dc.OnClose(func() { u.logger.Debug("[relay] tunnel DC closed") if u.mode == "dc" { u.closeAllConns() } }) dc.OnMessage(func(msg webrtc.DataChannelMessage) { u.modeOnce.Do(func() { u.mode = "dc" u.logger.Info("[relay] === MODE: DC ===") }) u.handleDCMessage(msg.Data) }) } sampleTrack, _ := webrtc.NewTrackLocalStaticSample( webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeVP8}, "video", "tunnel-video", ) u.sampleTrack = sampleTrack audioTrack, _ := webrtc.NewTrackLocalStaticRTP( webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeOpus}, "audio", "tunnel-audio", ) pc.AddTrack(audioTrack) pc.AddTrack(sampleTrack) ordered := true dcNotif, err := pc.CreateDataChannel("producerNotification", &webrtc.DataChannelInit{Ordered: &ordered}) if err == nil { dcNotif.OnOpen(func() { u.logger.Debug("[relay] producerNotification DC opened") }) dcNotif.OnMessage(func(msg webrtc.DataChannelMessage) { u.logger.Debug(fmt.Sprintf("[relay] producerNotification msg len=%d", len(msg.Data))) }) } dcCmd, err := pc.CreateDataChannel("producerCommand", &webrtc.DataChannelInit{Ordered: &ordered}) if err == nil { dcCmd.OnOpen(func() { u.logger.Debug("[relay] producerCommand DC opened") }) dcCmd.OnMessage(func(msg webrtc.DataChannelMessage) { u.logger.Debug(fmt.Sprintf("[relay] producerCommand msg len=%d", len(msg.Data))) }) } producerScreen, psErr := pc.CreateDataChannel("producerScreenShare", &webrtc.DataChannelInit{Ordered: &ordered}) if psErr == nil { u.producerScreen = producerScreen producerScreen.OnOpen(func() { u.logger.Debug("[relay] producerScreenShare DC open, reading uplink screen") }) producerScreen.OnMessage(func(msg webrtc.DataChannelMessage) { if u.sym != nil { u.sym.HandleScreenFrame(msg.Data) } }) } screenDC, scErr := pc.CreateDataChannel("consumerScreenShare", &webrtc.DataChannelInit{Ordered: &ordered}) if scErr == nil { u.screenDC = screenDC screenDC.OnOpen(func() { u.logger.Debug("[relay] consumerScreenShare DC open, writing downlink screen") }) } pc.OnICECandidate(func(cand *webrtc.ICECandidate) { if cand == nil { return } if u.externalICE != nil { u.externalICE(cand) } }) pc.OnConnectionStateChange(func(state webrtc.PeerConnectionState) { u.logger.Debug(fmt.Sprintf("[relay] connection state: %s (mode=%s)", state.String(), u.mode)) if u.externalCSC != nil { u.externalCSC(state) } }) pc.OnTrack(func(track *webrtc.TrackRemote, receiver *webrtc.RTPReceiver) { u.logger.Debug(fmt.Sprintf("[relay] remote track: %s", track.Codec().MimeType)) u.modeOnce.Do(func() { u.mode = "video" u.logger.Info("[relay] === MODE: VIDEO ===") u.tun = tunnel.NewVP8DataTunnel(sampleTrack, u.obf, u.logger) u.tun.Start(0, 0) var downlink tunnel.DataTunnel = u.tun if u.screenDC != nil { writer := tunnel.NewScreenWriter(u.obf, "screen-down", u.logger) dc := u.screenDC writer.SetSend(dc.Send) u.sym = tunnel.NewSymmetricScreenTunnel(u.tun, writer, u.obf, func() bool { return dc.ReadyState() == webrtc.DataChannelStateOpen }, u.logger) downlink = u.sym u.logger.Info("[relay] === MODE: VIDEO (with screenshare) ===") } if u.OnConnected != nil { u.OnConnected(downlink) } }) go u.readTrack(track) }) u.logger.Debug(fmt.Sprintf("[relay] PC created (%d ICE servers)", len(iceServers))) return nil } func (u *TunnelRelay) CreateOffer() (webrtc.SessionDescription, error) { offer, err := u.pc.CreateOffer(nil) if err != nil { return offer, err } u.pc.SetLocalDescription(offer) return offer, nil } func (u *TunnelRelay) CreateAnswer() (webrtc.SessionDescription, error) { answer, err := u.pc.CreateAnswer(nil) if err != nil { return answer, err } u.pc.SetLocalDescription(answer) return answer, nil } func (u *TunnelRelay) SetRemoteDescription(sdpType webrtc.SDPType, sdp string) error { err := u.pc.SetRemoteDescription(webrtc.SessionDescription{Type: sdpType, SDP: sdp}) if err != nil { return err } u.remoteSet = true for _, cand := range u.pending { u.pc.AddICECandidate(cand) } u.pending = nil return nil } func (u *TunnelRelay) AddICECandidate(candidate webrtc.ICECandidateInit) error { if !u.remoteSet { u.pending = append(u.pending, candidate) return nil } return u.pc.AddICECandidate(candidate) } func (u *TunnelRelay) OnICECandidate(fn func(*webrtc.ICECandidate)) { u.externalICE = fn } func (u *TunnelRelay) OnConnectionStateChange(fn func(webrtc.PeerConnectionState)) { u.externalCSC = fn } func (u *TunnelRelay) Close() { u.closeAllConns() if u.sym != nil { u.sym.Stop() u.sym = nil } if u.tun != nil { u.tun.Stop() u.tun = nil } u.dcMu.Lock() u.dc = nil u.dcMu.Unlock() if u.pc != nil { u.pc.OnConnectionStateChange(nil) u.pc.OnICECandidate(nil) u.pc.OnTrack(nil) oldPC := u.pc u.pc = nil go oldPC.Close() } u.remoteSet = false u.pending = nil u.sampleTrack = nil } func (u *TunnelRelay) handleDCMessage(data []byte) { if u.obf != nil { pt, ok := u.obf.DecryptPayload(data) if !ok { u.logger.Debug(fmt.Sprintf("[dc] decrypt failed, dropping %d bytes", len(data))) return } data = pt } if len(data) < 5 { return } connID := binary.BigEndian.Uint32(data[0:4]) mt := data[4] payload := data[5:] switch mt { case tunnel.MsgConnect: go u.connectTCP(connID, string(payload)) case tunnel.MsgUDP: go u.handleUDP(connID, payload) case tunnel.MsgData: val, ok := u.conns.Load(connID) if ok { dc := val.(*dcConn) cp := make([]byte, len(payload)) copy(cp, payload) select { case dc.ch <- cp: default: u.logger.Debug(fmt.Sprintf("[dc] conn %d write queue full, dropping %d bytes", connID, len(payload))) } } case tunnel.MsgClose: val, ok := u.conns.LoadAndDelete(connID) if ok { dc := val.(*dcConn) close(dc.ch) } } } func (u *TunnelRelay) sendDCFrame(connID uint32, mt byte, payload []byte) { u.dcMu.Lock() defer u.dcMu.Unlock() if u.dc == nil { return } buf := make([]byte, 5+len(payload)) binary.BigEndian.PutUint32(buf[0:4], connID) buf[4] = mt copy(buf[5:], payload) wire := buf if u.obf != nil { wire = u.obf.EncryptPayload(buf) if wire == nil { return } } u.dc.Send(wire) } func (u *TunnelRelay) connectTCP(connID uint32, addr string) { u.logger.Debug(fmt.Sprintf("[dc] CONNECT %d -> %s", connID, common.MaskAddr(addr))) conn, err := u.dialTCP(addr) if err != nil { u.logger.Warn(fmt.Sprintf("[dc] CONNECT %d failed: %s", connID, common.MaskError(err))) u.sendDCFrame(connID, tunnel.MsgConnectErr, []byte(common.MaskError(err))) return } dc := &dcConn{conn: conn, ch: make(chan []byte, 256)} u.conns.Store(connID, dc) u.sendDCFrame(connID, tunnel.MsgConnectOK, nil) u.logger.Debug(fmt.Sprintf("[dc] CONNECTED %d -> %s", connID, common.MaskAddr(addr))) go func() { for data := range dc.ch { conn.Write(data) } conn.Close() }() bufSz := u.readBufSize if bufSz <= 0 { bufSz = common.RTPBufSize } buf := make([]byte, bufSz) sent := 0 for { n, err := conn.Read(buf) if n > 0 { u.sendDCFrame(connID, tunnel.MsgData, buf[:n]) sent += n } if err != nil { if err != io.EOF { u.logger.Warn(fmt.Sprintf("[dc] conn %d read error: %s", connID, common.MaskError(err))) } break } } u.logger.Debug(fmt.Sprintf("[dc] conn %d closed, sent %d bytes", connID, sent)) u.sendDCFrame(connID, tunnel.MsgClose, nil) u.conns.Delete(connID) } func (u *TunnelRelay) handleUDP(connID uint32, payload []byte) { if len(payload) < 2 { return } addrLen := int(payload[0]) if len(payload) < 1+addrLen { return } addr := string(payload[1 : 1+addrLen]) data := payload[1+addrLen:] conn, err := u.dialUDP(addr) if err != nil { return } defer conn.Close() conn.SetDeadline(time.Now().Add(5 * time.Second)) conn.Write(data) resp := make([]byte, common.UDPBufSize) n, err := conn.Read(resp) if err != nil { return } u.sendDCFrame(connID, tunnel.MsgUDPReply, resp[:n]) } func (u *TunnelRelay) dialTCP(addr string) (net.Conn, error) { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() return u.dialer.DialContext(ctx, N.NetworkTCP, M.ParseSocksaddr(addr)) } func (u *TunnelRelay) dialUDP(addr string) (net.Conn, error) { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() return u.dialer.DialContext(ctx, N.NetworkUDP, M.ParseSocksaddr(addr)) } func (u *TunnelRelay) closeAllConns() { u.conns.Range(func(key, val any) bool { dc := val.(*dcConn) dc.conn.Close() u.conns.Delete(key) return true }) } func (u *TunnelRelay) readTrack(track *webrtc.TrackRemote) { if track.Codec().MimeType != webrtc.MimeTypeVP8 { buf := make([]byte, common.UDPBufSize) for { if _, _, err := track.Read(buf); err != nil { return } } } var vp8Pkt codecs.VP8Packet var pkt rtp.Packet var frameBuf []byte var lastSeq uint16 var haveLastSeq bool frameValid := false var recvCount int buf := make([]byte, common.RTPBufSize) for { n, _, err := track.Read(buf) if err != nil { return } if pkt.Unmarshal(buf[:n]) != nil { continue } if haveLastSeq && pkt.SequenceNumber != lastSeq+1 { frameValid = false frameBuf = frameBuf[:0] } lastSeq = pkt.SequenceNumber haveLastSeq = true vp8Payload, err := vp8Pkt.Unmarshal(pkt.Payload) if err != nil { frameValid = false frameBuf = frameBuf[:0] continue } if vp8Pkt.S == 1 { frameBuf = frameBuf[:0] frameValid = true } if !frameValid { continue } frameBuf = append(frameBuf, vp8Payload...) if !pkt.Marker { continue } recvCount++ if recvCount <= 3 || recvCount%200 == 0 { u.logger.Debug(fmt.Sprintf("[video] recv vp8 frame #%d %d bytes", recvCount, len(frameBuf))) } if u.tun != nil { u.tun.HandleFrame(frameBuf) } frameBuf = frameBuf[:0] frameValid = false } }