Add call protocol, Rmux. Update AmneziaWG. Fixes and improvements

This commit is contained in:
Shtorm
2026-08-06 14:19:34 +03:00
parent 051e928b01
commit d6b9f693c4
89 changed files with 15671 additions and 132 deletions

View File

@@ -0,0 +1,8 @@
package common
const (
UDPBufSize = 4096
RTPBufSize = 65536
VP8BufSize = 1126
DCBufSize = 32768
)

View File

@@ -0,0 +1,15 @@
package common
import (
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing/common/logger"
)
type ResolveFunc func(hostname string) (string, error)
type PeerConnectionConfigurer interface {
ConfigureSettingEngine(settingEngine *webrtc.SettingEngine)
}
type AddTunnelTracksFunc func(pc *webrtc.PeerConnection, logger logger.ContextLogger, prefix string) *webrtc.TrackLocalStaticSample
type ReadTrackFunc func(track *webrtc.TrackRemote, handler func([]byte), logger logger.ContextLogger, prefix string)

View File

@@ -0,0 +1,87 @@
package common
import (
"context"
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"os"
"strings"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
const UserAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/148.0.0.0 Safari/537.36"
func LoadCookies(path string) (string, error) {
data, err := os.ReadFile(path)
if err != nil {
return "", fmt.Errorf("cannot read cookies: %w", err)
}
var cookies []struct {
Name string `json:"name"`
Value string `json:"value"`
}
if err := json.Unmarshal(data, &cookies); err != nil {
return "", fmt.Errorf("cannot parse cookies: %w", err)
}
parts := make([]string, len(cookies))
for i, c := range cookies {
parts[i] = c.Name + "=" + c.Value
}
return strings.Join(parts, "; "), nil
}
func HttpClient(dialer N.Dialer) *http.Client {
return &http.Client{
Transport: &http.Transport{
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
return dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
},
},
}
}
func HttpGet(dialer N.Dialer, endpoint string) ([]byte, error) {
req, _ := http.NewRequest("GET", endpoint, nil)
req.Header.Set("User-Agent", UserAgent)
resp, err := HttpClient(dialer).Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
return io.ReadAll(resp.Body)
}
func CookieValue(cookieHeader, name string) string {
for _, part := range strings.Split(cookieHeader, ";") {
part = strings.TrimSpace(part)
eq := strings.IndexByte(part, '=')
if eq != -1 && part[:eq] == name {
return part[eq+1:]
}
}
return ""
}
func FilterCookies(cookieHeader string, allow []string) string {
allowed := make(map[string]struct{}, len(allow))
for _, n := range allow {
allowed[n] = struct{}{}
}
var out []string
for _, part := range strings.Split(cookieHeader, ";") {
trimmed := strings.TrimSpace(part)
eq := strings.IndexByte(trimmed, '=')
if eq == -1 {
continue
}
if _, ok := allowed[trimmed[:eq]]; ok {
out = append(out, trimmed)
}
}
return strings.Join(out, "; ")
}

View File

@@ -0,0 +1,58 @@
package common
import (
"net"
"strings"
)
func FixICEURL(iceURL string) string {
idx := strings.Index(iceURL, ":")
if idx < 0 {
return iceURL
}
scheme := iceURL[:idx]
if scheme != "turn" && scheme != "stun" && scheme != "turns" && scheme != "stuns" {
return iceURL
}
rest := iceURL[idx+1:]
if strings.HasPrefix(rest, "[") {
return iceURL
}
if strings.Count(rest, ":") <= 1 {
return iceURL
}
params := ""
if qm := strings.Index(rest, "?"); qm >= 0 {
params = rest[qm:]
rest = rest[:qm]
}
lastColon := strings.LastIndex(rest, ":")
if lastColon > 0 {
host := rest[:lastColon]
port := rest[lastColon+1:]
if net.ParseIP(host) != nil {
return scheme + ":[" + host + "]:" + port + params
}
}
if net.ParseIP(rest) != nil {
return scheme + ":[" + rest + "]" + params
}
return iceURL
}
func ExtractICEHost(iceURL string) string {
idx := strings.Index(iceURL, ":")
if idx < 0 {
return ""
}
rest := iceURL[idx+1:]
params := strings.Index(rest, "?")
if params >= 0 {
rest = rest[:params]
}
host, _, err := net.SplitHostPort(rest)
if err != nil {
return rest
}
return host
}

View File

@@ -0,0 +1,40 @@
package common
import (
"math/rand/v2"
"time"
)
const backoffJitterFloorDivisor = 4
func BackoffWithJitter(attempt int, initialDelay, maxDelay time.Duration) time.Duration {
if initialDelay <= 0 {
return 0
}
if maxDelay < initialDelay {
maxDelay = initialDelay
}
if attempt < 0 {
attempt = 0
}
ceiling := maxDelay
if shifted := initialDelay << uint(attempt); shifted > 0 && shifted < maxDelay {
ceiling = shifted
}
floor := ceiling / backoffJitterFloorDivisor
return floor + time.Duration(rand.Int64N(int64(ceiling-floor)+1))
}
func DurationInRange(minDuration, maxDuration time.Duration) time.Duration {
if maxDuration <= minDuration {
return minDuration
}
return minDuration + time.Duration(rand.Int64N(int64(maxDuration-minDuration)+1))
}
func IntInRange(minValue, maxValue int) int {
if maxValue <= minValue {
return minValue
}
return minValue + rand.IntN(maxValue-minValue+1)
}

View File

@@ -0,0 +1,68 @@
package common
import (
"fmt"
"net"
)
func MaskError(err error) string {
if err == nil {
return ""
}
if !MaskingEnabled {
return err.Error()
}
if opErr, ok := err.(*net.OpError); ok {
msg := opErr.Op
if opErr.Net != "" {
msg += " " + opErr.Net
}
if opErr.Source != nil {
msg += " " + MaskAddr(opErr.Source.String())
}
if opErr.Source != nil && opErr.Addr != nil {
msg += "->"
}
if opErr.Addr != nil {
msg += MaskAddr(opErr.Addr.String())
}
msg += ": " + opErr.Err.Error()
return msg
}
return err.Error()
}
const MaskingEnabled = true
func MaskAddr(addr string) string {
if !MaskingEnabled {
return addr
}
host, port, err := net.SplitHostPort(addr)
if err != nil {
host = addr
port = ""
}
masked := maskHost(host)
if port != "" {
return net.JoinHostPort(masked, port)
}
return masked
}
func maskHost(host string) string {
if host == "" {
return ""
}
ip := net.ParseIP(host)
if ip != nil {
if ip4 := ip.To4(); ip4 != nil {
return fmt.Sprintf("%d.%d.x.x", ip4[0], ip4[1])
}
return "x::x"
}
if len(host) <= 1 {
return "*"
}
return string(host[0]) + "***"
}

View File

@@ -0,0 +1,89 @@
package common
import (
"fmt"
"github.com/pion/rtp"
"github.com/pion/rtp/codecs"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing/common/logger"
)
func AddTunnelTracks(pc *webrtc.PeerConnection, logger logger.ContextLogger, prefix string) *webrtc.TrackLocalStaticSample {
sampleTrack, _ := webrtc.NewTrackLocalStaticSample(
webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeVP8},
"video", "tunnel-video",
)
audioTrack, _ := webrtc.NewTrackLocalStaticRTP(
webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeOpus},
"audio", "tunnel-audio",
)
audioSender, audioErr := pc.AddTrack(audioTrack)
videoSender, videoErr := pc.AddTrack(sampleTrack)
logger.Debug(fmt.Sprintf("%s: AddTrack audio: sender=%v err=%v", prefix, audioSender != nil, audioErr))
logger.Debug(fmt.Sprintf("%s: AddTrack video: sender=%v err=%v", prefix, videoSender != nil, videoErr))
logger.Debug(fmt.Sprintf("%s: senders count: %d", prefix, len(pc.GetSenders())))
return sampleTrack
}
func ReadTrack(track *webrtc.TrackRemote, handler func([]byte), logger logger.ContextLogger, prefix string) {
if track.Codec().MimeType != webrtc.MimeTypeVP8 {
buf := make([]byte, 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
recvCount := 0
buf := make([]byte, 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 {
logger.Debug(fmt.Sprintf("%s: recv vp8 frame #%d %d bytes", prefix, recvCount, len(frameBuf)))
}
if handler != nil {
frame := make([]byte, len(frameBuf))
copy(frame, frameBuf)
handler(frame)
}
frameBuf = frameBuf[:0]
frameValid = false
}
}

View File

@@ -0,0 +1,17 @@
package common
import (
"time"
"github.com/gorilla/websocket"
)
func CloseWS(ws *websocket.Conn) {
if ws == nil {
return
}
ws.WriteControl(websocket.CloseMessage,
websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""),
time.Now().Add(time.Second))
ws.Close()
}

158
transport/call/config.go Normal file
View File

@@ -0,0 +1,158 @@
package call
import (
"context"
"fmt"
"net"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/transport/call/dion"
"github.com/sagernet/sing-box/transport/call/telemost"
"github.com/sagernet/sing-box/transport/call/tunnel"
"github.com/sagernet/sing-box/transport/call/vk"
"github.com/sagernet/sing-box/transport/call/wbstream"
E "github.com/sagernet/sing/common/exceptions"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
type Role int
const (
RoleCreator Role = iota
RoleJoiner
)
type Config struct {
Platform string
Mode string
JoinLink string
Cookies string
CookieString string
Email string
Password string
ReadBuffer int
Role Role
Dialer N.Dialer
DNSRouter adapter.DNSRouter
Logger logger.ContextLogger
}
func Connect(ctx context.Context, cfg Config) (*Bridge, error) {
readBuf := cfg.ReadBuffer
if readBuf <= 0 {
readBuf = 32768
}
log := cfg.Logger
if log == nil {
log = logger.NOP()
}
cookieStr := cfg.CookieString
if cookieStr == "" {
cookieStr = cfg.Cookies
}
switch cfg.Platform {
case "telemost":
switch cfg.Role {
case RoleCreator:
relay, joinLink, err := telemost.ConnectCreator(ctx, cookieStr, cfg.JoinLink, readBuf, cfg.Dialer, log)
if err != nil {
return nil, err
}
log.Notice(fmt.Sprintf("call[telemost]: join_link=%s", joinLink))
return &Bridge{relay: relay}, nil
case RoleJoiner:
tun, err := telemost.ConnectJoiner(ctx, cfg.JoinLink, "", readBuf, cfg.Dialer, cfg.DNSRouter, log)
if err != nil {
return nil, err
}
relay := tunnel.NewRelayBridge(tun, "joiner", readBuf, cfg.Dialer, log)
relay.MarkReady()
return &Bridge{relay: relay}, nil
}
case "wbstream":
switch cfg.Role {
case RoleCreator:
relay, joinLink, err := wbstream.ConnectCreator(ctx, cookieStr, cfg.JoinLink, cfg.Mode, readBuf, cfg.Dialer, log)
if err != nil {
return nil, err
}
log.Notice(fmt.Sprintf("call[wbstream]: join_link=%s", joinLink))
return &Bridge{relay: relay}, nil
case RoleJoiner:
tun, err := wbstream.ConnectJoiner(ctx, cfg.JoinLink, "", cfg.Mode, readBuf, cfg.Dialer, cfg.DNSRouter, log)
if err != nil {
return nil, err
}
relay := tunnel.NewRelayBridge(tun, "joiner", readBuf, cfg.Dialer, log)
relay.MarkReady()
return &Bridge{relay: relay}, nil
}
case "vk":
switch cfg.Role {
case RoleCreator:
relay, joinLink, err := vk.ConnectCreator(ctx, cookieStr, cfg.JoinLink, readBuf, cfg.Dialer, log)
if err != nil {
return nil, err
}
log.Notice(fmt.Sprintf("call[vk]: join_link=%s", joinLink))
return &Bridge{relay: relay}, nil
case RoleJoiner:
tun, err := vk.ConnectJoiner(ctx, cfg.JoinLink, "", readBuf, cfg.Dialer, cfg.DNSRouter, log)
if err != nil {
return nil, err
}
relay := tunnel.NewRelayBridge(tun, "joiner", readBuf, cfg.Dialer, log)
relay.MarkReady()
return &Bridge{relay: relay}, nil
}
case "dion":
switch cfg.Role {
case RoleCreator:
relay, joinLink, err := dion.ConnectCreator(ctx, cookieStr, cfg.JoinLink, cfg.Email, cfg.Password, readBuf, cfg.Dialer, log)
if err != nil {
return nil, err
}
log.Notice(fmt.Sprintf("call[dion]: join_link=%s", joinLink))
return &Bridge{relay: relay}, nil
case RoleJoiner:
tun, err := dion.ConnectJoiner(ctx, cfg.JoinLink, "", readBuf, cfg.Dialer, log)
if err != nil {
return nil, err
}
relay := tunnel.NewRelayBridge(tun, "joiner", readBuf, cfg.Dialer, log)
relay.MarkReady()
return &Bridge{relay: relay}, nil
}
}
return nil, E.New("call: unsupported platform ", cfg.Platform)
}
type Bridge struct {
relay *tunnel.RelayBridge
}
func NewBridge(relay *tunnel.RelayBridge) *Bridge {
return &Bridge{relay: relay}
}
func (b *Bridge) Close() error {
b.relay.Close()
return nil
}
func (b *Bridge) DialContext(ctx context.Context, destination string) (net.Conn, error) {
return b.relay.DialContext(ctx, destination)
}
func (b *Bridge) ListenPacket(ctx context.Context, destination string) (net.Conn, error) {
return b.relay.ListenPacket(ctx, destination)
}
func (b *Bridge) SetAcceptHandler(fn func(conn net.Conn, destination string)) {
b.relay.SetAcceptHandler(fn)
}
func (b *Bridge) SetUDPAcceptHandler(fn func(conn net.Conn, destination string)) {
b.relay.SetUDPAcceptHandler(fn)
}

650
transport/call/dion/api.go Normal file
View File

@@ -0,0 +1,650 @@
package dion
import (
"bytes"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/http/cookiejar"
"strings"
"sync"
"time"
"github.com/google/uuid"
"github.com/sagernet/sing-box/transport/call/common"
N "github.com/sagernet/sing/common/network"
)
var ErrSessionExpired = errors.New("dion: session expired, re-login required")
var errLoginEndpointMissing = errors.New("dion: login endpoint not available")
const (
accessCookieName = "vc-access-token"
refreshCookieName = "vc-refresh-token"
loginClientsPath = "/v2/users/login/web"
loginPlatformPath = "/platform/v2/auth/auth-providers/dion/login/password"
)
const (
refreshSkewSeconds = 60
refreshMaxAttempts = 3
refreshBaseDelay = 2 * time.Second
refreshDelayMultiply = 1.75
)
const (
APIBase = "https://api.dion.vc"
APIClientsBase = "https://api-clients.dion.vc"
WebBase = "https://dion.vc"
Origin = "https://dion.vc"
CookieDomain = "dion.vc"
)
type GuestUser struct {
ID string `json:"id"`
Name string `json:"name"`
Email string `json:"email"`
Initials string `json:"initials"`
Position string `json:"position"`
AvatarHTTPPath string `json:"avatar_http_path"`
IsProfileFilledIn bool `json:"is_profile_filled_in"`
Roles []string `json:"roles"`
}
type GuestAuthResponse struct {
AccessToken string `json:"access_token"`
AuthProvider string `json:"auth_provider"`
IsAuthBySSO bool `json:"is_auth_by_sso"`
User GuestUser `json:"user"`
}
type LoginResponse struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
AuthProvider string `json:"auth_provider"`
IsAuthBySSO bool `json:"is_auth_by_sso"`
User GuestUser `json:"user"`
}
type EventInfo struct {
ID string `json:"id"`
Name string `json:"name"`
Slug string `json:"slug"`
OrgID string `json:"org_id"`
Admins []string `json:"admins"`
PSTN struct {
Number string `json:"number"`
Pin int `json:"pin"`
Prefix string `json:"prefix"`
} `json:"pstn"`
}
type WSSConnectResponse struct {
Host string `json:"host"`
Path string `json:"path"`
Schema string `json:"schema"`
URL string `json:"url"`
Params map[string]string `json:"params"`
}
type Session struct {
HTTPClient *http.Client
Device DeviceProfile
AccessToken string
AccessTokenExp time.Time
UserID string
SessionID string
cookiesPath string
email string
password string
refreshMu sync.Mutex
}
type AuthResult struct {
Session *Session
Event *EventInfo
WSS *WSSConnectResponse
SessionID string
}
func NewSession(dialer N.Dialer) (*Session, error) {
jar, err := cookiejar.New(nil)
if err != nil {
return nil, fmt.Errorf("cookiejar: %w", err)
}
httpClient := common.HttpClient(dialer)
httpClient.Jar = jar
return &Session{HTTPClient: httpClient, Device: RandomDeviceProfile()}, nil
}
func (s *Session) RegisterGuest() (*GuestAuthResponse, error) {
auth, err := s.callRefreshOnce()
if err != nil {
return nil, err
}
s.applyRefreshResult(auth)
return auth, nil
}
func (s *Session) RegisterAnonymousGuest(eventID, displayName string) (*GuestAuthResponse, error) {
if eventID == "" {
return nil, fmt.Errorf("empty event_id")
}
if displayName == "" {
displayName = "Guest"
}
body, _ := json.Marshal(map[string]any{
"event_id": eventID,
"name": displayName,
})
req, err := http.NewRequest(http.MethodPost, APIBase+"/platform/v1/users/register/guest", bytes.NewReader(body))
if err != nil {
return nil, err
}
s.setBaseHeaders(req, "")
req.Header.Set("Content-Type", "application/json")
resp, err := s.HTTPClient.Do(req)
if err != nil {
return nil, fmt.Errorf("register/guest: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated {
return nil, fmt.Errorf("register/guest: status %d: %s", resp.StatusCode, string(raw))
}
var auth GuestAuthResponse
if err := json.Unmarshal(raw, &auth); err != nil {
return nil, fmt.Errorf("register/guest decode: %w", err)
}
if auth.AccessToken == "" {
return nil, fmt.Errorf("register/guest: empty access_token: %s", string(raw))
}
s.applyRefreshResult(&auth)
return &auth, nil
}
// SetCredentials stores an email/password pair used by refreshLocked to
// re-authenticate when the refresh cookie is missing or rejected.
func (s *Session) SetCredentials(email, password string) {
s.email = strings.TrimSpace(email)
s.password = password
}
// LoginWithPassword exchanges credentials for a fresh token pair. The web
// front-end posts to api-clients, and switches to the platform endpoint when
// the DION_PLATFORM_COOKIE_AUTH_ENABLED toggle is on, so both are tried.
func (s *Session) LoginWithPassword(email, password string) error {
if email == "" || password == "" {
return fmt.Errorf("login: email and password are required")
}
body, _ := json.Marshal(map[string]string{"email": email, "password": password})
login, err := s.postLogin(APIClientsBase+loginClientsPath, body)
if errors.Is(err, errLoginEndpointMissing) {
login, err = s.postLogin(APIBase+loginPlatformPath, body)
}
if err != nil {
return err
}
s.email = email
s.password = password
s.applyLoginResult(login)
return nil
}
func (s *Session) postLogin(target string, body []byte) (*LoginResponse, error) {
req, err := http.NewRequest(http.MethodPost, target, bytes.NewReader(body))
if err != nil {
return nil, err
}
s.setBaseHeaders(req, "")
req.Header.Set("Content-Type", "application/json")
resp, err := s.HTTPClient.Do(req)
if err != nil {
return nil, fmt.Errorf("login: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusMethodNotAllowed {
return nil, errLoginEndpointMissing
}
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated {
return nil, fmt.Errorf("login: status %d: %s", resp.StatusCode, string(raw))
}
var login LoginResponse
if err := json.Unmarshal(raw, &login); err != nil {
return nil, fmt.Errorf("login decode: %w", err)
}
if login.AccessToken == "" {
return nil, fmt.Errorf("login: empty access_token: %s", string(raw))
}
return &login, nil
}
func (s *Session) applyLoginResult(login *LoginResponse) {
s.AccessToken = login.AccessToken
s.UserID = login.User.ID
if exp, err := parseJWTExpiry(login.AccessToken); err == nil {
s.AccessTokenExp = exp
}
if login.RefreshToken != "" {
s.SetCookieInJar(refreshCookieName, login.RefreshToken)
}
s.SetCookieInJar(accessCookieName, login.AccessToken)
}
func (s *Session) Refresh() error {
s.refreshMu.Lock()
defer s.refreshMu.Unlock()
return s.refreshLocked()
}
func (s *Session) EnsureValidToken() error {
s.refreshMu.Lock()
defer s.refreshMu.Unlock()
if s.AccessToken != "" && !s.AccessTokenExp.IsZero() &&
time.Until(s.AccessTokenExp) > time.Duration(refreshSkewSeconds)*time.Second {
return nil
}
return s.refreshLocked()
}
func (s *Session) DoAuthenticated(buildRequest func() (*http.Request, error)) (*http.Response, error) {
if err := s.EnsureValidToken(); err != nil {
return nil, err
}
req, err := buildRequest()
if err != nil {
return nil, err
}
s.setBaseHeaders(req, s.AccessToken)
resp, err := s.HTTPClient.Do(req)
if err != nil {
return nil, err
}
if resp.StatusCode != http.StatusUnauthorized {
return resp, nil
}
staleToken := s.AccessToken
io.Copy(io.Discard, resp.Body)
resp.Body.Close()
s.refreshMu.Lock()
if s.AccessToken == staleToken {
if err := s.refreshLocked(); err != nil {
s.refreshMu.Unlock()
return nil, err
}
}
s.refreshMu.Unlock()
retryReq, err := buildRequest()
if err != nil {
return nil, err
}
s.setBaseHeaders(retryReq, s.AccessToken)
return s.HTTPClient.Do(retryReq)
}
func (s *Session) WhoAmI() (json.RawMessage, error) {
resp, err := s.DoAuthenticated(func() (*http.Request, error) {
return http.NewRequest(http.MethodGet, APIBase+"/platform/v1/whoami", nil)
})
if err != nil {
return nil, fmt.Errorf("whoami: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("whoami: status %d: %s", resp.StatusCode, string(raw))
}
return raw, nil
}
func (s *Session) GetEventBySlug(slug string) (*EventInfo, error) {
if slug == "" {
return nil, fmt.Errorf("empty room ID")
}
eventURL := fmt.Sprintf("%s/conference/v1/events/slug/%s", APIBase, slug)
resp, err := s.DoAuthenticated(func() (*http.Request, error) {
req, err := http.NewRequest(http.MethodGet, eventURL, nil)
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
return req, nil
})
if err != nil {
return nil, fmt.Errorf("get event: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("get event: status %d: %s", resp.StatusCode, string(raw))
}
var event EventInfo
if err := json.Unmarshal(raw, &event); err != nil {
return nil, fmt.Errorf("get event decode: %w", err)
}
if event.ID == "" {
return nil, fmt.Errorf("get event: empty id: %s", string(raw))
}
return &event, nil
}
func (s *Session) GenerateSlug() (string, error) {
resp, err := s.DoAuthenticated(func() (*http.Request, error) {
return http.NewRequest(http.MethodGet, APIClientsBase+"/v2/events/slug/generate", nil)
})
if err != nil {
return "", fmt.Errorf("generate room ID: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("generate room ID: status %d: %s", resp.StatusCode, string(raw))
}
var out struct {
Slug string `json:"slug"`
}
if err := json.Unmarshal(raw, &out); err != nil {
return "", fmt.Errorf("generate room ID decode: %w", err)
}
if out.Slug == "" {
return "", fmt.Errorf("generate room ID: empty: %s", string(raw))
}
return out.Slug, nil
}
type CreateEventOptions struct {
Slug string
EventParams []string
IsImpersonalSlug bool
IsOnCloud bool
}
func (s *Session) CreateEvent(opts CreateEventOptions) (*EventInfo, error) {
if opts.Slug == "" {
return nil, fmt.Errorf("empty room ID")
}
if opts.EventParams == nil {
opts.EventParams = []string{"guest_access"}
}
body, _ := json.Marshal(map[string]any{
"event_params": opts.EventParams,
"is_impersonal_slug": opts.IsImpersonalSlug,
"is_on_cloud": opts.IsOnCloud,
"slug": opts.Slug,
})
resp, err := s.DoAuthenticated(func() (*http.Request, error) {
req, err := http.NewRequest(http.MethodPost, APIBase+"/conference/v1/events", bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
return req, nil
})
if err != nil {
return nil, fmt.Errorf("create event: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated {
return nil, fmt.Errorf("create event: status %d: %s", resp.StatusCode, string(raw))
}
var event EventInfo
if err := json.Unmarshal(raw, &event); err != nil {
return nil, fmt.Errorf("create event decode: %w", err)
}
if event.ID == "" {
return nil, fmt.Errorf("create event: empty id: %s", string(raw))
}
return &event, nil
}
func (s *Session) CreateRoom() (*EventInfo, error) {
slug, err := s.GenerateSlug()
if err != nil {
return nil, err
}
return s.CreateEvent(CreateEventOptions{
Slug: slug,
EventParams: []string{"guest_access"},
IsImpersonalSlug: true,
IsOnCloud: true,
})
}
func (s *Session) ConnectWSS(sessionID string) (*WSSConnectResponse, error) {
if sessionID == "" {
sessionID = uuid.New().String()
}
body, _ := json.Marshal(map[string]string{"session_id": sessionID})
resp, err := s.DoAuthenticated(func() (*http.Request, error) {
req, err := http.NewRequest(http.MethodPost, APIBase+"/conference/v1/connect/wss", bytes.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/json")
return req, nil
})
if err != nil {
return nil, fmt.Errorf("connect/wss: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("connect/wss: status %d: %s", resp.StatusCode, string(raw))
}
var wss WSSConnectResponse
if err := json.Unmarshal(raw, &wss); err != nil {
return nil, fmt.Errorf("connect/wss decode: %w", err)
}
if wss.URL == "" {
return nil, fmt.Errorf("connect/wss: empty url: %s", string(raw))
}
s.SessionID = sessionID
return &wss, nil
}
func (s *Session) LookupEventBySlugAnonymous(slug string) (*EventInfo, error) {
if slug == "" {
return nil, fmt.Errorf("empty room ID")
}
eventURL := fmt.Sprintf("%s/conference/v1/events/slug/%s", APIBase, slug)
req, err := http.NewRequest(http.MethodGet, eventURL, nil)
if err != nil {
return nil, err
}
s.setBaseHeaders(req, "")
req.Header.Set("Content-Type", "application/json")
resp, err := s.HTTPClient.Do(req)
if err != nil {
return nil, fmt.Errorf("get event anon: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("get event anon: status %d: %s", resp.StatusCode, string(raw))
}
var event EventInfo
if err := json.Unmarshal(raw, &event); err != nil {
return nil, fmt.Errorf("get event anon decode: %w", err)
}
if event.ID == "" {
return nil, fmt.Errorf("get event anon: empty id: %s", string(raw))
}
return &event, nil
}
func JoinAsGuest(dialer N.Dialer, slug, displayName string) (*Session, *EventInfo, error) {
session, err := NewSession(dialer)
if err != nil {
return nil, nil, err
}
if err := session.PrimeCookies(slug); err != nil {
return nil, nil, fmt.Errorf("prime cookies: %w", err)
}
event, err := session.LookupEventBySlugAnonymous(slug)
if err != nil {
return nil, nil, err
}
if _, err := session.RegisterAnonymousGuest(event.ID, displayName); err != nil {
return nil, nil, fmt.Errorf("RegisterAnonymousGuest: %w", err)
}
return session, event, nil
}
func AuthAndGetTicket(dialer N.Dialer, slug string) (*AuthResult, error) {
session, err := NewSession(dialer)
if err != nil {
return nil, err
}
if err := session.PrimeCookies(slug); err != nil {
return nil, err
}
if _, err := session.RegisterGuest(); err != nil {
return nil, err
}
if _, err := session.WhoAmI(); err != nil {
return nil, fmt.Errorf("whoami after guest auth: %w", err)
}
event, err := session.GetEventBySlug(slug)
if err != nil {
return nil, err
}
sessionID := uuid.New().String()
wss, err := session.ConnectWSS(sessionID)
if err != nil {
return nil, err
}
return &AuthResult{
Session: session,
Event: event,
WSS: wss,
SessionID: sessionID,
}, nil
}
func ParseRoom(input string) string {
trimmed := strings.TrimSpace(input)
if trimmed == "" {
return ""
}
trimmed = strings.TrimPrefix(trimmed, "dion://")
trimmed = strings.TrimPrefix(trimmed, "https://")
trimmed = strings.TrimPrefix(trimmed, "http://")
trimmed = strings.TrimPrefix(trimmed, "dion.vc/")
trimmed = strings.TrimPrefix(trimmed, "event/")
if idx := strings.Index(trimmed, "?"); idx >= 0 {
trimmed = trimmed[:idx]
}
if idx := strings.Index(trimmed, "/"); idx >= 0 {
trimmed = trimmed[:idx]
}
return trimmed
}
func (s *Session) setBaseHeaders(req *http.Request, accessToken string) {
req.Header.Set("User-Agent", s.Device.UserAgent)
req.Header.Set("Origin", Origin)
req.Header.Set("Referer", Origin+"/")
req.Header.Set("Accept", "*/*")
req.Header.Set("Accept-Language", "en")
req.Header.Set("X-Request-Id", uuid.New().String())
for name, value := range s.Device.Headers() {
req.Header.Set(name, value)
}
if accessToken != "" {
req.Header.Set("Authorization", "Bearer "+accessToken)
}
}
func (s *Session) callRefreshOnce() (*GuestAuthResponse, error) {
req, err := http.NewRequest(http.MethodPost, APIBase+"/platform/v1/auth/refresh/web", bytes.NewReader(nil))
if err != nil {
return nil, err
}
s.setBaseHeaders(req, "")
req.Header.Set("Content-Length", "0")
resp, err := s.HTTPClient.Do(req)
if err != nil {
return nil, fmt.Errorf("auth/refresh/web: %w", err)
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode >= 400 && resp.StatusCode < 500 {
return nil, fmt.Errorf("%w: status %d: %s", ErrSessionExpired, resp.StatusCode, string(raw))
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("auth/refresh/web: status %d: %s", resp.StatusCode, string(raw))
}
var auth GuestAuthResponse
if err := json.Unmarshal(raw, &auth); err != nil {
return nil, fmt.Errorf("auth/refresh/web decode: %w", err)
}
if auth.AccessToken == "" {
return nil, fmt.Errorf("auth/refresh/web: empty access_token: %s", string(raw))
}
return &auth, nil
}
func (s *Session) applyRefreshResult(auth *GuestAuthResponse) {
s.AccessToken = auth.AccessToken
s.UserID = auth.User.ID
if exp, err := parseJWTExpiry(auth.AccessToken); err == nil {
s.AccessTokenExp = exp
}
s.SetCookieInJar(accessCookieName, auth.AccessToken)
}
func (s *Session) refreshLocked() error {
if !s.HasRefreshCookie() && s.HasCredentials() {
return s.LoginWithPassword(s.email, s.password)
}
var lastErr error
delay := refreshBaseDelay
for attempt := 1; attempt <= refreshMaxAttempts; attempt++ {
auth, err := s.callRefreshOnce()
if err == nil {
s.applyRefreshResult(auth)
return nil
}
if errors.Is(err, ErrSessionExpired) {
if s.HasCredentials() {
return s.LoginWithPassword(s.email, s.password)
}
return err
}
lastErr = err
if attempt < refreshMaxAttempts {
time.Sleep(delay)
delay = time.Duration(float64(delay) * refreshDelayMultiply)
}
}
return fmt.Errorf("refresh failed after %d attempts: %w", refreshMaxAttempts, lastErr)
}
func parseJWTExpiry(token string) (time.Time, error) {
parts := strings.Split(token, ".")
if len(parts) < 2 {
return time.Time{}, fmt.Errorf("invalid JWT")
}
payload, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil {
return time.Time{}, fmt.Errorf("decode payload: %w", err)
}
var claims struct {
Exp int64 `json:"exp"`
}
if err := json.Unmarshal(payload, &claims); err != nil {
return time.Time{}, fmt.Errorf("parse claims: %w", err)
}
if claims.Exp == 0 {
return time.Time{}, fmt.Errorf("no exp claim")
}
return time.Unix(claims.Exp, 0), nil
}

732
transport/call/dion/call.go Normal file
View File

@@ -0,0 +1,732 @@
package dion
import (
"fmt"
"sync"
"sync/atomic"
"time"
"github.com/google/uuid"
"github.com/pion/rtp/codecs"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
"github.com/sagernet/sing-box/transport/call/tunnel"
)
const (
sendVideoMidIndex = 12
sendScreenShareMidIndex = 13
recvScreenShareMidStr = "14"
defaultRecvVideoMid = "1"
recvVideoMidCount = 9
creatorVP8FPS = 24
creatorVP8Batch = 10
joinerVP8FPS = 24
joinerVP8Batch = 5
)
type Role string
const (
RoleCreator Role = "creator"
RoleJoiner Role = "joiner"
)
type CallConfig struct {
Auth *Session
Event *EventInfo
Obfuscator *tunnel.TunnelObfuscator
DisplayName string
Logger logger.ContextLogger
RecvMid string
Role Role
SettingEngine *webrtc.SettingEngine
Dialer N.Dialer
DNSRouter adapter.DNSRouter
}
type PeerEntry struct {
SessionID string
UserID string
Name string
CamState bool
JoinedAt time.Time
}
type Call struct {
cfg CallConfig
signaling *SignalingClient
peer *PionPeer
sendTrack *webrtc.TrackLocalStaticSample
vp8tun *tunnel.VP8DataTunnel
mySessionID string
peersMu sync.Mutex
peersByID map[string]*PeerEntry
subscribed map[string]bool
peerToMid map[string]string
freeMids []string
pendingSubs []string
onConnectedFired atomic.Bool
OnConnected func(tunnel.DataTunnel)
OnPeerRestart func()
OnRemoteSDP func(sdp string)
done chan struct{}
closeOnce sync.Once
}
func NewCall(cfg CallConfig) *Call {
if cfg.Logger == nil {
cfg.Logger = logger.NOP()
}
if cfg.Role == "" {
cfg.Role = RoleCreator
}
if cfg.RecvMid == "" {
cfg.RecvMid = defaultRecvVideoMid
}
var freeMids []string
if cfg.Role == RoleJoiner {
freeMids = []string{recvScreenShareMidStr}
} else {
freeMids = make([]string, 0, recvVideoMidCount)
for midIndex := 1; midIndex < recvVideoMidCount; midIndex++ {
freeMids = append(freeMids, fmt.Sprintf("%d", midIndex))
}
freeMids = append(freeMids, "0")
}
return &Call{
cfg: cfg,
peersByID: make(map[string]*PeerEntry),
subscribed: make(map[string]bool),
peerToMid: make(map[string]string),
freeMids: freeMids,
done: make(chan struct{}),
}
}
func (c *Call) Done() <-chan struct{} { return c.done }
func (c *Call) SessionID() string { return c.mySessionID }
func (c *Call) Start() error {
sessionID := uuid.New().String()
c.mySessionID = sessionID
c.cfg.Logger.Debug(fmt.Sprintf("[call] my session_id=%s", sessionID))
wss, err := c.cfg.Auth.ConnectWSS(sessionID)
if err != nil {
return fmt.Errorf("ConnectWSS: %w", err)
}
signaling, err := DialSignaling(wss.URL, SignalingDialOptions{
UserAgent: c.cfg.Auth.Device.UserAgent,
Logger: c.cfg.Logger,
Dialer: c.cfg.Dialer,
})
if err != nil {
return fmt.Errorf("DialSignaling: %w", err)
}
c.signaling = signaling
if err := signaling.WaitConnected(15 * time.Second); err != nil {
return fmt.Errorf("WaitConnected: %w", err)
}
youJoinedChan := make(chan YouJoinedParams, 1)
sdpAnswerChan := make(chan SDPAnswerParams, 4)
var onceYouJoined sync.Once
signaling.OnYouJoined = func(params YouJoinedParams) {
onceYouJoined.Do(func() { youJoinedChan <- params })
}
signaling.OnSDPAnswer = func(answerSDP string, transceivers []TransceiverDesc) {
select {
case sdpAnswerChan <- SDPAnswerParams{Answer: answerSDP, Transceivers: transceivers}:
default:
}
}
signaling.OnSpeakerJoined = c.handleSpeakerJoined
signaling.OnSpeakerDisconnected = c.handleSpeakerDisconnected
signaling.OnSpeakerCamStateChanged = c.handleSpeakerCamStateChanged
signaling.OnConfSpeakersState = c.handleConfSpeakersState
signaling.OnGetVideoFromUserResponse = c.handleGetVideoFromUserResponse
signaling.OnGetScreenSharingFromUserResponse = c.handleGetScreenSharingFromUserResponse
readLoopDone := make(chan error, 1)
go func() { readLoopDone <- signaling.ReadLoop() }()
if err := signaling.Subscribe(c.cfg.Event.ID, sessionID); err != nil {
return fmt.Errorf("Subscribe: %w", err)
}
var youJoined YouJoinedParams
select {
case youJoined = <-youJoinedChan:
case err := <-readLoopDone:
return fmt.Errorf("read loop ended before you_joined: %v", err)
case <-time.After(15 * time.Second):
return fmt.Errorf("timeout waiting for you_joined")
}
c.cfg.Logger.Debug(fmt.Sprintf("[call] you_joined ice_servers=%d", len(youJoined.IceServers)))
pionAPI := NewPionAPI(c.cfg.SettingEngine)
iceServers := ResolveICEServerHosts(youJoined.IceServers, c.cfg.DNSRouter, c.cfg.Dialer, c.cfg.Logger)
peer, err := BuildPionPeer(pionAPI, iceServers)
if err != nil {
return fmt.Errorf("BuildPionPeer: %w", err)
}
c.peer = peer
sendMidIndex := sendVideoMidIndex
trackLabel := "dion-tunnel-" + sessionID
if c.cfg.Role == RoleCreator {
sendMidIndex = sendScreenShareMidIndex
trackLabel = "dion-tunnel-screen-" + sessionID
}
track, err := webrtc.NewTrackLocalStaticSample(
webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeVP8, ClockRate: 90000},
"video", trackLabel,
)
if err != nil {
return fmt.Errorf("NewTrackLocalStaticSample: %w", err)
}
c.sendTrack = track
if len(peer.Transceivers) <= sendMidIndex {
return fmt.Errorf("transceiver layout short, have %d", len(peer.Transceivers))
}
sender := peer.Transceivers[sendMidIndex].Sender()
if sender == nil {
return fmt.Errorf("mid=%d sender nil", sendMidIndex)
}
if err := sender.ReplaceTrack(track); err != nil {
return fmt.Errorf("ReplaceTrack: %w", err)
}
c.cfg.Logger.Debug(fmt.Sprintf("[call] role=%s attached send track to mid=%d", c.cfg.Role, sendMidIndex))
peer.PC.OnTrack(func(remoteTrack *webrtc.TrackRemote, _ *webrtc.RTPReceiver) {
c.cfg.Logger.Debug(fmt.Sprintf("[call] OnTrack id=%q kind=%s codec=%s ssrc=%d",
remoteTrack.ID(), remoteTrack.Kind().String(), remoteTrack.Codec().MimeType, remoteTrack.SSRC()))
if remoteTrack.Codec().MimeType != webrtc.MimeTypeVP8 {
go drainTrack(remoteTrack)
return
}
go c.readVP8Track(remoteTrack)
})
var pendingMu sync.Mutex
pendingCandidates := make([]webrtc.ICECandidateInit, 0, 32)
remoteSet := false
sendCandidate := func(cand webrtc.ICECandidateInit) {
entry := ICECandidateJSON{Candidate: cand.Candidate}
if cand.SDPMid != nil {
m := *cand.SDPMid
entry.SDPMid = &m
}
if cand.SDPMLineIndex != nil {
i := *cand.SDPMLineIndex
entry.SDPMLineIndex = &i
}
if cand.UsernameFragment != nil {
entry.UsernameFragment = *cand.UsernameFragment
}
if err := signaling.SendICECandidates([]ICECandidateJSON{entry}); err != nil {
c.cfg.Logger.Warn(fmt.Sprintf("[ice] SendICECandidates: %v", err))
}
}
flushPending := func() {
pendingMu.Lock()
toFlush := pendingCandidates
pendingCandidates = nil
pendingMu.Unlock()
for _, cand := range toFlush {
sendCandidate(cand)
}
}
peer.PC.OnICECandidate(func(cand *webrtc.ICECandidate) {
if cand == nil {
return
}
init := cand.ToJSON()
pendingMu.Lock()
alreadyRemote := remoteSet
if !alreadyRemote {
pendingCandidates = append(pendingCandidates, init)
}
pendingMu.Unlock()
if alreadyRemote {
sendCandidate(init)
}
})
iceConnected := make(chan struct{}, 1)
iceDead := make(chan webrtc.ICEConnectionState, 1)
peer.PC.OnICEConnectionStateChange(func(state webrtc.ICEConnectionState) {
c.cfg.Logger.Debug(fmt.Sprintf("[ice] state=%s", state.String()))
switch state {
case webrtc.ICEConnectionStateConnected, webrtc.ICEConnectionStateCompleted:
select {
case iceConnected <- struct{}{}:
default:
}
case webrtc.ICEConnectionStateFailed, webrtc.ICEConnectionStateClosed:
select {
case iceDead <- state:
default:
}
}
})
envelope, _, err := peer.CreateAndSetOffer()
if err != nil {
return fmt.Errorf("CreateAndSetOffer: %w", err)
}
offerParams := SDPOfferParams{
MicState: false,
CamState: false,
NoiseSuppressionState: true,
ScreenSharingQuality: "default",
Datachannels: peer.DatachannelDescs,
Transceivers: peer.TransceiverDescs,
Offer: envelope,
}
if err := signaling.SendSDPOffer(offerParams); err != nil {
return fmt.Errorf("SendSDPOffer: %w", err)
}
var answer SDPAnswerParams
select {
case answer = <-sdpAnswerChan:
case err := <-readLoopDone:
return fmt.Errorf("read loop ended before sdp_answer: %v", err)
case <-time.After(20 * time.Second):
return fmt.Errorf("timeout waiting for sdp_answer")
}
if c.OnRemoteSDP != nil {
c.OnRemoteSDP(answer.Answer)
}
if err := peer.PC.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeAnswer,
SDP: answer.Answer,
}); err != nil {
return fmt.Errorf("SetRemoteDescription: %w", err)
}
pendingMu.Lock()
remoteSet = true
pendingMu.Unlock()
flushPending()
select {
case <-iceConnected:
case state := <-iceDead:
return fmt.Errorf("ICE died before connected: %s", state.String())
case err := <-readLoopDone:
return fmt.Errorf("read loop ended before ICE connected: %v", err)
case <-time.After(30 * time.Second):
return fmt.Errorf("timeout waiting for ICE connected; state=%s", peer.PC.ICEConnectionState().String())
}
c.cfg.Logger.Debug("[ice] connected")
fps, batch := joinerVP8FPS, joinerVP8Batch
if c.cfg.Role == RoleCreator {
fps, batch = creatorVP8FPS, creatorVP8Batch
}
c.vp8tun = tunnel.NewVP8DataTunnel(c.sendTrack, c.cfg.Obfuscator, c.cfg.Logger)
c.vp8tun.Start(fps, batch)
c.fireOnConnected(c.vp8tun)
if c.cfg.Role == RoleCreator {
if err := signaling.SendScreenSharingSwitchOn(); err != nil {
c.cfg.Logger.Warn(fmt.Sprintf("[call] SendScreenSharingSwitchOn: %v", err))
} else {
c.cfg.Logger.Debug("[call] sent screensharing_switch_on")
}
if err := signaling.SendScreensharingQualityChange("good"); err != nil {
c.cfg.Logger.Warn(fmt.Sprintf("[call] SendScreensharingQualityChange: %v", err))
} else {
c.cfg.Logger.Debug("[call] sent screensharing_quality_change=good")
}
} else {
if err := signaling.SendCamStateChange(true); err != nil {
c.cfg.Logger.Warn(fmt.Sprintf("[call] SendCamStateChange: %v", err))
} else {
c.cfg.Logger.Debug("[call] sent cam_state_change=true")
}
}
go c.discoverPeersAndSubscribe()
go c.runStatReporter()
go func() {
defer close(c.done)
select {
case state := <-iceDead:
c.cfg.Logger.Debug(fmt.Sprintf("[call] ICE went to %s", state.String()))
case err := <-readLoopDone:
c.cfg.Logger.Debug(fmt.Sprintf("[call] read loop ended: %v", err))
}
}()
return nil
}
func (c *Call) Close() {
c.closeOnce.Do(func() {
if c.vp8tun != nil {
c.vp8tun.Stop()
}
if c.signaling != nil {
c.signaling.Close()
}
if c.peer != nil {
c.peer.Close()
}
})
}
func (c *Call) fireOnConnected(tun tunnel.DataTunnel) {
if !c.onConnectedFired.CompareAndSwap(false, true) {
return
}
if c.OnConnected != nil {
c.OnConnected(tun)
}
}
func (c *Call) handleSpeakerJoined(params SpeakerJoinedParams) {
if params.SessionID == c.mySessionID {
return
}
c.peersMu.Lock()
_, wasKnown := c.peersByID[params.SessionID]
c.peersByID[params.SessionID] = &PeerEntry{
SessionID: params.SessionID,
UserID: params.UserID,
Name: params.Name,
CamState: params.CamState,
JoinedAt: time.Now(),
}
var toKick []string
if !wasKnown {
for id := range c.peersByID {
if id != params.SessionID {
toKick = append(toKick, id)
}
}
}
c.peersMu.Unlock()
c.cfg.Logger.Debug(fmt.Sprintf("[call] speaker_joined session_id=%s name=%q cam=%v", params.SessionID, params.Name, params.CamState))
for _, staleID := range toKick {
if err := c.signaling.SendKickOne(staleID); err != nil {
c.cfg.Logger.Debug(fmt.Sprintf("[call] SendKickOne(%s): %v", staleID, err))
continue
}
c.cfg.Logger.Debug(fmt.Sprintf("[call] kicked stale peer session_id=%s for newcomer=%s", staleID, params.SessionID))
c.peersMu.Lock()
delete(c.peersByID, staleID)
delete(c.subscribed, staleID)
c.releaseMidLocked(staleID)
c.peersMu.Unlock()
}
if len(toKick) > 0 && c.OnPeerRestart != nil {
c.cfg.Logger.Debug(fmt.Sprintf("[call] firing OnPeerRestart from kick path (kicked=%d newcomer=%s)", len(toKick), params.SessionID))
c.OnPeerRestart()
}
if c.cfg.Role == RoleJoiner || params.CamState {
c.subscribeIfNeeded(params.SessionID)
}
}
func (c *Call) handleSpeakerDisconnected(params SpeakerDisconnectedParams) {
c.peersMu.Lock()
delete(c.peersByID, params.SessionID)
delete(c.subscribed, params.SessionID)
c.releaseMidLocked(params.SessionID)
var freshestUnsubscribed string
var freshestAt time.Time
for sid, entry := range c.peersByID {
if c.subscribed[sid] {
continue
}
if entry.JoinedAt.After(freshestAt) {
freshestAt = entry.JoinedAt
freshestUnsubscribed = sid
}
}
hasFreeMid := len(c.freeMids) > 0
c.peersMu.Unlock()
c.cfg.Logger.Debug(fmt.Sprintf("[call] speaker_disconnected session_id=%s", params.SessionID))
if freshestUnsubscribed != "" && hasFreeMid {
c.cfg.Logger.Debug(fmt.Sprintf("[call] claiming freed mid for unsubscribed peer %s", freshestUnsubscribed))
c.subscribeIfNeeded(freshestUnsubscribed)
}
}
func (c *Call) handleSpeakerCamStateChanged(params SpeakerCamStateChangedParams) {
if params.SessionID == c.mySessionID {
return
}
c.peersMu.Lock()
if entry, ok := c.peersByID[params.SessionID]; ok {
entry.CamState = params.CamState
} else {
c.peersByID[params.SessionID] = &PeerEntry{SessionID: params.SessionID, CamState: params.CamState, JoinedAt: time.Now()}
}
c.peersMu.Unlock()
c.cfg.Logger.Debug(fmt.Sprintf("[call] speaker_cam_state_changed session_id=%s cam=%v", params.SessionID, params.CamState))
if params.CamState {
c.subscribeIfNeeded(params.SessionID)
}
}
func (c *Call) handleConfSpeakersState(response ConfSpeakersStateResponse) {
for _, entry := range response.Speakers {
if entry.SessionID == c.mySessionID || entry.SessionID == "" {
continue
}
c.peersMu.Lock()
c.peersByID[entry.SessionID] = &PeerEntry{
SessionID: entry.SessionID,
UserID: entry.UserID,
Name: entry.Name,
CamState: entry.CamState,
JoinedAt: time.Now(),
}
c.peersMu.Unlock()
if c.cfg.Role == RoleJoiner || entry.CamState {
c.subscribeIfNeeded(entry.SessionID)
}
}
}
func (c *Call) discoverPeersAndSubscribe() {
time.Sleep(500 * time.Millisecond)
if err := c.signaling.SendConfSpeakersState(DefaultConfSpeakersStateRequest()); err != nil {
c.cfg.Logger.Debug(fmt.Sprintf("[call] SendConfSpeakersState: %v", err))
}
}
func (c *Call) subscribeIfNeeded(peerSessionID string) {
c.peersMu.Lock()
if c.subscribed[peerSessionID] {
c.peersMu.Unlock()
return
}
entry := c.peersByID[peerSessionID]
if entry == nil {
c.peersMu.Unlock()
return
}
if len(c.freeMids) == 0 {
c.peersMu.Unlock()
c.cfg.Logger.Debug(fmt.Sprintf("[call] no free recv mid for peer %s, ignoring", peerSessionID))
return
}
mid := c.freeMids[0]
c.freeMids = c.freeMids[1:]
c.peerToMid[peerSessionID] = mid
c.subscribed[peerSessionID] = true
c.pendingSubs = append(c.pendingSubs, peerSessionID)
c.peersMu.Unlock()
var sendErr error
if c.cfg.Role == RoleJoiner {
sendErr = c.signaling.SendGetScreenSharingFromUser(GetScreenSharingFromUserRequest{
SessionID: entry.SessionID,
TransceiverID: mid,
UserID: entry.UserID,
})
} else {
sendErr = c.signaling.SendGetVideoFromUser(GetVideoFromUserRequest{
SessionID: entry.SessionID,
TransceiverID: mid,
UserID: entry.UserID,
Username: entry.Name,
})
}
if sendErr != nil {
c.cfg.Logger.Debug(fmt.Sprintf("[call] subscribe to peer %s failed: %v", peerSessionID, sendErr))
c.peersMu.Lock()
delete(c.subscribed, peerSessionID)
delete(c.peerToMid, peerSessionID)
c.freeMids = append(c.freeMids, mid)
if len(c.pendingSubs) > 0 && c.pendingSubs[len(c.pendingSubs)-1] == peerSessionID {
c.pendingSubs = c.pendingSubs[:len(c.pendingSubs)-1]
}
c.peersMu.Unlock()
return
}
c.cfg.Logger.Debug(fmt.Sprintf("[call] subscribed to %s on mid=%s", peerSessionID, mid))
if c.OnPeerRestart != nil {
c.cfg.Logger.Debug(fmt.Sprintf("[call] firing OnPeerRestart from subscribe path (peer=%s)", peerSessionID))
c.OnPeerRestart()
}
}
func (c *Call) handleGetVideoFromUserResponse(resp GetVideoFromUserResponse, errCode int, errMsg string) {
c.handleSubscribeResponse("get_video_from_user", resp.SessionID, resp.TransceiverID, errCode, errMsg)
}
func (c *Call) handleGetScreenSharingFromUserResponse(resp GetScreenSharingFromUserResponse, errCode int, errMsg string) {
c.handleSubscribeResponse("get_screensharing_from_user", resp.SessionID, resp.TransceiverID, errCode, errMsg)
}
func (c *Call) handleSubscribeResponse(rpc, sessionID, transceiverID string, errCode int, errMsg string) {
c.peersMu.Lock()
if sessionID == "" && len(c.pendingSubs) > 0 {
sessionID = c.pendingSubs[0]
c.pendingSubs = c.pendingSubs[1:]
} else if len(c.pendingSubs) > 0 {
for i, pending := range c.pendingSubs {
if pending == sessionID {
c.pendingSubs = append(c.pendingSubs[:i], c.pendingSubs[i+1:]...)
break
}
}
}
if errCode != 0 {
mid := c.peerToMid[sessionID]
delete(c.subscribed, sessionID)
delete(c.peerToMid, sessionID)
if mid != "" {
c.freeMids = append(c.freeMids, mid)
}
c.peersMu.Unlock()
c.cfg.Logger.Debug(fmt.Sprintf("[call] %s FAILED session=%s mid=%s code=%d msg=%q", rpc, sessionID, mid, errCode, errMsg))
return
}
c.peersMu.Unlock()
c.cfg.Logger.Debug(fmt.Sprintf("[call] %s OK session=%s mid=%s", rpc, sessionID, transceiverID))
}
func (c *Call) releaseMidLocked(peerSessionID string) {
if mid, ok := c.peerToMid[peerSessionID]; ok {
delete(c.peerToMid, peerSessionID)
c.freeMids = append(c.freeMids, mid)
}
}
func (c *Call) readVP8Track(track *webrtc.TrackRemote) {
var vp8Pkt codecs.VP8Packet
var frameBuf []byte
var lastSeq uint16
var haveLastSeq bool
frameValid := false
for {
pkt, _, err := track.ReadRTP()
if err != nil {
return
}
if pkt == 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
}
if c.vp8tun != nil {
c.vp8tun.HandleFrame(frameBuf)
}
frameBuf = frameBuf[:0]
frameValid = false
}
}
func (c *Call) runStatReporter() {
select {
case <-c.done:
return
case <-time.After(1500 * time.Millisecond):
}
if err := c.signaling.SendPCIceStat(); err != nil {
c.cfg.Logger.Debug(fmt.Sprintf("[call] SendPCIceStat: %v", err))
} else {
c.cfg.Logger.Debug("[call] sent pc_ice_stat")
}
c.sendStatReport()
ticker := time.NewTicker(10 * time.Second)
defer ticker.Stop()
for {
select {
case <-c.done:
return
case <-ticker.C:
c.sendStatReport()
}
}
}
func (c *Call) sendStatReport() {
report := ClientStatReport{
ReportTimeUnixMS: time.Now().UnixMilli(),
Connection: ClientStatConnection{},
}
report.Audio.In = ClientStatAudioIn{Codec: "opus", IsEnabled: true, Mid: 9}
report.Video.In = c.buildVideoInStats()
outStat := ClientStatVideoOut{
Mid: sendVideoMidIndex,
Codec: "VP8",
IsEnabled: true,
Resolution: ClientStatResolution{Width: 1280, Height: 720},
Framerate: c.vp8tun.FPS(),
ScalabilityMode: "L1T1",
}
report.Video.Out = outStat
report.Video.OutV2 = []ClientStatVideoOut{outStat}
if err := c.signaling.SendClientStatZip(report); err != nil {
c.cfg.Logger.Debug(fmt.Sprintf("[call] SendClientStatZip: %v", err))
}
}
func (c *Call) buildVideoInStats() []ClientStatVideoIn {
c.peersMu.Lock()
defer c.peersMu.Unlock()
out := make([]ClientStatVideoIn, 0, len(c.peerToMid))
for sessionID, midStr := range c.peerToMid {
midInt := 0
fmt.Sscanf(midStr, "%d", &midInt)
out = append(out, ClientStatVideoIn{
Codec: "VP8",
IsEnabled: true,
Mid: midInt,
Resolution: ClientStatResolution{Width: 1280, Height: 720},
Framerate: c.vp8tun.FPS(),
SessionID: sessionID,
})
}
return out
}
func drainTrack(track *webrtc.TrackRemote) {
buf := make([]byte, 1500)
for {
if _, _, err := track.Read(buf); err != nil {
return
}
}
}

View File

@@ -0,0 +1,139 @@
package dion
import (
"context"
"fmt"
"time"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing-box/transport/call/tunnel"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
func ConnectCreator(ctx context.Context, cookieStr, roomID, email, password string, readBuf int, dialer N.Dialer, logger logger.ContextLogger) (*tunnel.RelayBridge, string, error) {
auth, err := NewSession(dialer)
if err != nil {
return nil, "", fmt.Errorf("dion: new session: %w", err)
}
if err := auth.LoadCookieString(cookieStr); err != nil {
return nil, "", fmt.Errorf("dion: load cookies: %w", err)
}
auth.SetCredentials(email, password)
if err := auth.EnsureValidToken(); err != nil {
return nil, "", fmt.Errorf("dion: ensure valid token: %w", err)
}
requestedRoom := ParseRoom(roomID)
var event *EventInfo
if requestedRoom != "" {
event, err = auth.GetEventBySlug(requestedRoom)
if err != nil {
return nil, "", fmt.Errorf("dion: get event by slug: %w", err)
}
} else {
event, err = auth.CreateRoom()
if err != nil {
return nil, "", fmt.Errorf("dion: create room: %w", err)
}
}
joinLink := WebBase + "/event/" + event.Slug
if readBuf <= 0 {
readBuf = 32768
}
obf, err := tunnel.NewTunnelObfuscator(tunnel.DeriveSecretFromJoinLink(event.Slug))
if err != nil {
return nil, "", fmt.Errorf("dion: obfuscator init: %w", err)
}
relayCh := make(chan *tunnel.RelayBridge, 1)
var activeRelay *tunnel.RelayBridge
call := NewCall(CallConfig{
Auth: auth,
Event: event,
Obfuscator: obf,
DisplayName: "Creator",
Logger: logger,
Dialer: dialer,
Role: RoleCreator,
})
call.OnConnected = func(tun tunnel.DataTunnel) {
if activeRelay != nil {
activeRelay.Reset()
}
bridgeReadBuf := common.VP8BufSize
if _, ok := tun.(*tunnel.DCTunnel); ok {
bridgeReadBuf = readBuf
}
activeRelay = tunnel.NewRelayBridge(tun, "creator", bridgeReadBuf, dialer, logger)
activeRelay.MarkReady()
select {
case relayCh <- activeRelay:
default:
}
}
call.OnPeerRestart = func() {
if activeRelay != nil {
activeRelay.Reset()
}
}
go func() {
if err := call.Start(); err != nil {
logger.Error(fmt.Sprintf("dion: call start failed: %v", err))
}
}()
select {
case relay := <-relayCh:
return relay, joinLink, nil
case <-ctx.Done():
call.Close()
return nil, "", ctx.Err()
case <-time.After(60 * time.Second):
call.Close()
return nil, "", fmt.Errorf("dion: creator tunnel timed out")
}
}
func ConnectJoiner(ctx context.Context, roomID, displayName string, readBuf int, dialer N.Dialer, logger logger.ContextLogger) (tunnel.DataTunnel, error) {
if displayName == "" {
displayName = "Joiner"
}
slug := ParseRoom(roomID)
if slug == "" {
return nil, fmt.Errorf("dion: missing room")
}
auth, event, err := JoinAsGuest(dialer, slug, displayName)
if err != nil {
return nil, fmt.Errorf("dion: join as guest: %w", err)
}
obf, err := tunnel.NewTunnelObfuscator(tunnel.DeriveSecretFromJoinLink(event.Slug))
if err != nil {
return nil, fmt.Errorf("dion: obfuscator init: %w", err)
}
call := NewCall(CallConfig{
Auth: auth,
Event: event,
Obfuscator: obf,
DisplayName: displayName,
Logger: logger,
Dialer: dialer,
Role: RoleJoiner,
})
tunCh := make(chan tunnel.DataTunnel, 1)
call.OnConnected = func(tun tunnel.DataTunnel) {
select {
case tunCh <- tun:
default:
}
}
go func() {
if err := call.Start(); err != nil {
logger.Error(fmt.Sprintf("dion: call start failed: %v", err))
}
}()
select {
case tun := <-tunCh:
return tun, nil
case <-ctx.Done():
call.Close()
return nil, ctx.Err()
}
}

View File

@@ -0,0 +1,134 @@
package dion
import (
"fmt"
"io"
"net/http"
"net/http/cookiejar"
"net/url"
"strings"
)
type CookieEntry struct {
Name string `json:"name"`
Value string `json:"value"`
}
func (s *Session) LoadCookies(entries []CookieEntry) error {
if s.HTTPClient.Jar == nil {
jar, err := cookiejar.New(nil)
if err != nil {
return fmt.Errorf("cookiejar: %w", err)
}
s.HTTPClient.Jar = jar
}
web, err := url.Parse(WebBase)
if err != nil {
return fmt.Errorf("parse target %s: %w", WebBase, err)
}
cookies := make([]*http.Cookie, 0, len(entries))
for _, entry := range entries {
if entry.Name == "" {
continue
}
cookies = append(cookies, &http.Cookie{
Name: entry.Name,
Value: entry.Value,
Path: "/",
Domain: CookieDomain,
})
}
s.HTTPClient.Jar.SetCookies(web, cookies)
s.seedAccessTokenFromCookies(entries)
return nil
}
func (s *Session) LoadCookieString(cookieStr string) error {
cookieStr = strings.TrimSpace(cookieStr)
if cookieStr == "" {
return fmt.Errorf("empty cookie string")
}
var entries []CookieEntry
for _, piece := range strings.Split(cookieStr, ";") {
piece = strings.TrimSpace(piece)
if piece == "" {
continue
}
eq := strings.IndexByte(piece, '=')
if eq <= 0 {
continue
}
entries = append(entries, CookieEntry{Name: piece[:eq], Value: piece[eq+1:]})
}
return s.LoadCookies(entries)
}
func (s *Session) SetCookieInJar(name, value string) {
if s.HTTPClient == nil || s.HTTPClient.Jar == nil {
return
}
web, err := url.Parse(WebBase)
if err != nil {
return
}
s.HTTPClient.Jar.SetCookies(web, []*http.Cookie{
{Name: name, Value: value, Path: "/", Domain: CookieDomain},
})
}
func (s *Session) HasCredentials() bool {
return s.email != "" && s.password != ""
}
func (s *Session) HasRefreshCookie() bool {
if s.HTTPClient == nil || s.HTTPClient.Jar == nil {
return false
}
web, err := url.Parse(WebBase)
if err != nil {
return false
}
for _, c := range s.HTTPClient.Jar.Cookies(web) {
if c.Name == refreshCookieName && c.Value != "" {
return true
}
}
return false
}
func (s *Session) PrimeCookies(slug string) error {
target := WebBase + "/"
if slug != "" {
target = fmt.Sprintf("%s/event/%s?showWeb=true", WebBase, slug)
}
req, err := http.NewRequest(http.MethodGet, target, nil)
if err != nil {
return err
}
s.setBaseHeaders(req, "")
resp, err := s.HTTPClient.Do(req)
if err != nil {
return fmt.Errorf("prime cookies: %w", err)
}
defer resp.Body.Close()
io.Copy(io.Discard, resp.Body)
if resp.StatusCode >= 400 {
return fmt.Errorf("prime cookies: status %d", resp.StatusCode)
}
return nil
}
func (s *Session) seedAccessTokenFromCookies(entries []CookieEntry) {
for _, entry := range entries {
if entry.Name != accessCookieName || entry.Value == "" {
continue
}
exp, err := parseJWTExpiry(entry.Value)
if err != nil {
return
}
s.AccessToken = entry.Value
s.AccessTokenExp = exp
return
}
}

View File

@@ -0,0 +1,185 @@
package dion
import (
"fmt"
"math/rand/v2"
)
type DeviceProfile struct {
UserAgent string
Platform string
BrowserType string
BrowserVersion string
DeviceBrand string
DeviceModel string
DeviceType string
OS string
OSVersion string
ScreenWidth int
ScreenHeight int
}
type deviceTemplate struct {
os string
osVersionPool []string
deviceBrandPool []string
deviceModelPool []string
browsers []browserTemplate
}
type browserTemplate struct {
browserType string
versionPool []string
userAgentFn func(osVersion, browserVersion string) string
}
var commonScreens = [][2]int{
{1280, 720}, {1366, 768}, {1440, 900}, {1536, 864},
{1600, 900}, {1680, 1050}, {1728, 1117}, {1920, 1080},
{2048, 1152}, {2560, 1440}, {2880, 1800}, {3840, 2160},
}
var deviceTemplates = []deviceTemplate{
{
os: "Mac OS",
osVersionPool: []string{"10.15.7", "11.7.10", "12.7.6", "13.6.9", "14.6.1", "15.1.0"},
deviceBrandPool: []string{"Apple"},
deviceModelPool: []string{"Macintosh"},
browsers: []browserTemplate{
{
browserType: "Chrome",
versionPool: []string{"141.0.0.0", "143.0.0.0", "145.0.0.0", "147.0.0.0", "149.0.0.0"},
userAgentFn: func(osVersion, browserVersion string) string {
return fmt.Sprintf("Mozilla/5.0 (Macintosh; Intel Mac OS X %s) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/%s Safari/537.36",
macOSVersionForUA(osVersion), browserVersion)
},
},
{
browserType: "Safari",
versionPool: []string{"17.6", "18.0", "18.1", "18.2"},
userAgentFn: func(osVersion, browserVersion string) string {
return fmt.Sprintf("Mozilla/5.0 (Macintosh; Intel Mac OS X %s) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/%s Safari/605.1.15",
macOSVersionForUA(osVersion), browserVersion)
},
},
{
browserType: "Firefox",
versionPool: []string{"128.0", "131.0", "133.0", "135.0"},
userAgentFn: func(osVersion, browserVersion string) string {
return fmt.Sprintf("Mozilla/5.0 (Macintosh; Intel Mac OS X %s; rv:%s) Gecko/20100101 Firefox/%s",
macOSVersionForUA(osVersion), browserVersion, browserVersion)
},
},
},
},
{
os: "Windows",
osVersionPool: []string{"10", "11"},
deviceBrandPool: []string{"Dell", "Lenovo", "HP", "Asus", "Acer", "MSI"},
deviceModelPool: []string{"PC"},
browsers: []browserTemplate{
{
browserType: "Chrome",
versionPool: []string{"141.0.0.0", "143.0.0.0", "145.0.0.0", "147.0.0.0", "149.0.0.0"},
userAgentFn: func(osVersion, browserVersion string) string {
return fmt.Sprintf("Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/%s Safari/537.36",
browserVersion)
},
},
{
browserType: "Edge",
versionPool: []string{"141.0.0.0", "143.0.0.0", "145.0.0.0", "147.0.0.0"},
userAgentFn: func(osVersion, browserVersion string) string {
return fmt.Sprintf("Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/%s Safari/537.36 Edg/%s",
browserVersion, browserVersion)
},
},
{
browserType: "Firefox",
versionPool: []string{"128.0", "131.0", "133.0", "135.0"},
userAgentFn: func(osVersion, browserVersion string) string {
return fmt.Sprintf("Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:%s) Gecko/20100101 Firefox/%s",
browserVersion, browserVersion)
},
},
},
},
{
os: "Linux",
osVersionPool: []string{"x86_64", "x86_64 GNU"},
deviceBrandPool: []string{"Dell", "Lenovo", "HP", "System76", "Framework"},
deviceModelPool: []string{"PC"},
browsers: []browserTemplate{
{
browserType: "Chrome",
versionPool: []string{"141.0.0.0", "143.0.0.0", "145.0.0.0", "147.0.0.0", "149.0.0.0"},
userAgentFn: func(osVersion, browserVersion string) string {
return fmt.Sprintf("Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/%s Safari/537.36",
browserVersion)
},
},
{
browserType: "Firefox",
versionPool: []string{"128.0", "131.0", "133.0", "135.0"},
userAgentFn: func(osVersion, browserVersion string) string {
return fmt.Sprintf("Mozilla/5.0 (X11; Linux x86_64; rv:%s) Gecko/20100101 Firefox/%s",
browserVersion, browserVersion)
},
},
},
},
}
func RandomDeviceProfile() DeviceProfile {
tmpl := deviceTemplates[rand.IntN(len(deviceTemplates))]
browser := tmpl.browsers[rand.IntN(len(tmpl.browsers))]
osVersion := tmpl.osVersionPool[rand.IntN(len(tmpl.osVersionPool))]
browserVersion := browser.versionPool[rand.IntN(len(browser.versionPool))]
screen := commonScreens[rand.IntN(len(commonScreens))]
return DeviceProfile{
UserAgent: browser.userAgentFn(osVersion, browserVersion),
Platform: "web",
BrowserType: browser.browserType,
BrowserVersion: browserVersion,
DeviceBrand: tmpl.deviceBrandPool[rand.IntN(len(tmpl.deviceBrandPool))],
DeviceModel: tmpl.deviceModelPool[rand.IntN(len(tmpl.deviceModelPool))],
DeviceType: "pc",
OS: tmpl.os,
OSVersion: osVersion,
ScreenWidth: screen[0],
ScreenHeight: screen[1],
}
}
func (p DeviceProfile) Headers() map[string]string {
return map[string]string{
"d-platform": p.Platform,
"d-browser-type": p.BrowserType,
"d-browser-version": p.BrowserVersion,
"d-device-brand": p.DeviceBrand,
"d-device-model": p.DeviceModel,
"d-device-type": p.DeviceType,
"d-os": p.OS,
"d-os-version": p.OSVersion,
"d-screen-height": fmt.Sprintf("%d", p.ScreenHeight),
"d-screen-width": fmt.Sprintf("%d", p.ScreenWidth),
}
}
func macOSVersionForUA(osVersion string) string {
switch osVersion {
case "10.15.7":
return "10_15_7"
case "11.7.10":
return "10_15_7"
case "12.7.6":
return "10_15_7"
case "13.6.9":
return "10_15_7"
case "14.6.1":
return "10_15_7"
case "15.1.0":
return "10_15_7"
}
return "10_15_7"
}

View File

@@ -0,0 +1,262 @@
package dion
import (
"context"
"fmt"
"net"
"net/netip"
"strings"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/dialer"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
type TransceiverPlan struct {
Mid int
Direction webrtc.RTPTransceiverDirection
Kind webrtc.RTPCodecType
Ctype string
}
type DataChannelPlan struct {
ID uint16
Label string
}
var DionTransceiverLayout = []TransceiverPlan{
{Mid: 0, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 1, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 2, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 3, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 4, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 5, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 6, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 7, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 8, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 9, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeAudio, Ctype: "Audio"},
{Mid: 10, Direction: webrtc.RTPTransceiverDirectionSendonly, Kind: webrtc.RTPCodecTypeAudio, Ctype: "Audio"},
{Mid: 11, Direction: webrtc.RTPTransceiverDirectionSendonly, Kind: webrtc.RTPCodecTypeAudio, Ctype: "AudioScreenSharing"},
{Mid: 12, Direction: webrtc.RTPTransceiverDirectionSendonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 13, Direction: webrtc.RTPTransceiverDirectionSendonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "ScreenSharing"},
{Mid: 14, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "ScreenSharing"},
{Mid: 15, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Video"},
{Mid: 16, Direction: webrtc.RTPTransceiverDirectionRecvonly, Kind: webrtc.RTPCodecTypeVideo, Ctype: "Padding"},
}
var DionDataChannels = []DataChannelPlan{
{ID: 0, Label: "vad"},
{ID: 1, Label: "stats"},
{ID: 2, Label: "speed"},
{ID: 3, Label: "video_quality"},
{ID: 4, Label: "media_messages"},
}
type PionPeer struct {
PC *webrtc.PeerConnection
Transceivers []*webrtc.RTPTransceiver
DataChannels map[string]*webrtc.DataChannel
TransceiverDescs []TransceiverDesc
DatachannelDescs []DataChannelDesc
}
func NewPionAPI(customEngine ...*webrtc.SettingEngine) *webrtc.API {
mediaEngine := &webrtc.MediaEngine{}
if err := mediaEngine.RegisterDefaultCodecs(); err != nil {
panic(fmt.Errorf("dion: register default codecs: %w", err))
}
engine := webrtc.SettingEngine{}
if len(customEngine) > 0 && customEngine[0] != nil {
engine = *customEngine[0]
}
return webrtc.NewAPI(
webrtc.WithMediaEngine(mediaEngine),
webrtc.WithSettingEngine(engine),
)
}
func ResolveICEServerHosts(entries []ICEServerEntry, dnsRouter adapter.DNSRouter, d N.Dialer, logger logger.ContextLogger) []ICEServerEntry {
if dnsRouter == nil {
return entries
}
resolved := make(map[string]string)
out := make([]ICEServerEntry, 0, len(entries))
for _, entry := range entries {
urls := make([]string, len(entry.URLs))
copy(urls, entry.URLs)
for k, raw := range urls {
host := extractICEHost(raw)
if host == "" {
continue
}
ip, ok := resolved[host]
if !ok {
var addrs []netip.Addr
var err error
addrs, err = dnsRouter.Lookup(context.Background(), host, d.(dialer.ResolveDialer).QueryOptions())
if err != nil {
logger.Warn(fmt.Sprintf("[dion] resolve ICE host %s failed: %v", host, err))
continue
}
ip = addrs[0].String()
resolved[host] = ip
logger.Debug(fmt.Sprintf("[dion] resolved ICE host %s -> %s", host, addrs[0]))
}
urls[k] = strings.Replace(raw, host, ip, 1)
}
out = append(out, ICEServerEntry{URLs: urls, Username: entry.Username, Credential: entry.Credential})
}
return out
}
func IceServerEntriesToWebRTC(entries []ICEServerEntry) []webrtc.ICEServer {
out := make([]webrtc.ICEServer, 0, len(entries))
for _, entry := range entries {
out = append(out, webrtc.ICEServer{
URLs: entry.URLs,
Username: entry.Username,
Credential: entry.Credential,
})
}
return out
}
func BuildPionPeer(api *webrtc.API, iceServers []ICEServerEntry) (*PionPeer, error) {
pc, err := api.NewPeerConnection(webrtc.Configuration{
ICEServers: IceServerEntriesToWebRTC(iceServers),
BundlePolicy: webrtc.BundlePolicyMaxBundle,
RTCPMuxPolicy: webrtc.RTCPMuxPolicyRequire,
})
if err != nil {
return nil, fmt.Errorf("new peer connection: %w", err)
}
transceivers := make([]*webrtc.RTPTransceiver, 0, len(DionTransceiverLayout))
for _, plan := range DionTransceiverLayout {
transceiver, err := pc.AddTransceiverFromKind(plan.Kind, webrtc.RTPTransceiverInit{
Direction: plan.Direction,
})
if err != nil {
pc.Close()
return nil, fmt.Errorf("add transceiver mid=%d kind=%s dir=%s: %w", plan.Mid, plan.Kind, plan.Direction, err)
}
transceivers = append(transceivers, transceiver)
}
dataChannels := make(map[string]*webrtc.DataChannel, len(DionDataChannels))
for _, plan := range DionDataChannels {
negotiated := true
id := plan.ID
dc, err := pc.CreateDataChannel(plan.Label, &webrtc.DataChannelInit{
Negotiated: &negotiated,
ID: &id,
})
if err != nil {
pc.Close()
return nil, fmt.Errorf("create datachannel %s id=%d: %w", plan.Label, plan.ID, err)
}
dataChannels[plan.Label] = dc
}
dcDescs := make([]DataChannelDesc, 0, len(DionDataChannels))
for _, plan := range DionDataChannels {
dcDescs = append(dcDescs, DataChannelDesc{ID: int(plan.ID), Label: plan.Label})
}
return &PionPeer{
PC: pc,
Transceivers: transceivers,
DataChannels: dataChannels,
DatachannelDescs: dcDescs,
}, nil
}
func (p *PionPeer) BuildOfferDescriptors() error {
if len(p.Transceivers) != len(DionTransceiverLayout) {
return fmt.Errorf("transceiver count drift: have %d want %d", len(p.Transceivers), len(DionTransceiverLayout))
}
descs := make([]TransceiverDesc, 0, len(p.Transceivers))
for index, transceiver := range p.Transceivers {
mid := transceiver.Mid()
if mid == "" {
return fmt.Errorf("transceiver index=%d has empty mid; call SetLocalDescription first", index)
}
plan := DionTransceiverLayout[index]
descs = append(descs, TransceiverDesc{
TransceiverID: mid,
SessionID: "",
Direction: directionToDion(plan.Direction),
Ctype: plan.Ctype,
})
}
p.TransceiverDescs = descs
return nil
}
func (p *PionPeer) CreateAndSetOffer() (offerEnvelope string, sdpOffer string, err error) {
offer, err := p.PC.CreateOffer(nil)
if err != nil {
return "", "", fmt.Errorf("create offer: %w", err)
}
if err := p.PC.SetLocalDescription(offer); err != nil {
return "", "", fmt.Errorf("set local description: %w", err)
}
if err := p.BuildOfferDescriptors(); err != nil {
return "", "", err
}
envelope, err := BuildSDPOfferEnvelope(offer.SDP)
if err != nil {
return "", "", fmt.Errorf("build envelope: %w", err)
}
return envelope, offer.SDP, nil
}
func (p *PionPeer) ApplyAnswerEnvelope(answerEnvelope string) error {
sdp, err := DecodeSDPAnswerInner(answerEnvelope)
if err != nil {
return fmt.Errorf("decode answer envelope: %w", err)
}
return p.PC.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeAnswer,
SDP: sdp,
})
}
func (p *PionPeer) Close() error {
if p.PC == nil {
return nil
}
return p.PC.Close()
}
func directionToDion(direction webrtc.RTPTransceiverDirection) string {
switch direction {
case webrtc.RTPTransceiverDirectionSendonly:
return "SendOnly"
case webrtc.RTPTransceiverDirectionRecvonly:
return "RecvOnly"
case webrtc.RTPTransceiverDirectionSendrecv:
return "SendRecv"
case webrtc.RTPTransceiverDirectionInactive:
return "Inactive"
}
return "Unknown"
}
func extractICEHost(raw string) string {
value := raw
for _, prefix := range []string{"stun:", "turn:", "turns:"} {
value = strings.TrimPrefix(value, prefix)
}
if idx := strings.Index(value, "?"); idx >= 0 {
value = value[:idx]
}
if idx := strings.LastIndex(value, ":"); idx >= 0 {
value = value[:idx]
}
if value == "" {
return ""
}
if net.ParseIP(value) != nil {
return ""
}
return value
}

View File

@@ -0,0 +1,692 @@
package dion
import (
"context"
"encoding/base64"
"encoding/json"
"fmt"
"net"
"net/http"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/gorilla/websocket"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
const (
MethodServerConnected = "server:notify:main:connected"
MethodServerYouJoined = "server:you_joined"
MethodServerSubscribeResponse = "server:main:response:subscribe:conference"
MethodServerSDPAnswer = "server:sdp_answer"
MethodServerSpeakerJoined = "server:speaker_joined"
MethodServerSpeakerDisconnected = "server:speaker_disconnected"
MethodServerHeartbeat = "server:notify:main:heartbeat"
MethodServerSpeakersResponse = "server:response:speakers"
MethodServerSpeakersResponseZip = "server:response:speakers_zip"
MethodClientSubscribeConference = "client:main:request:subscribe:conference"
MethodClientSDPOffer = "client:request:media:sdp_offer"
MethodClientSendICECandidates = "client:request:send_ice_candidates_zip"
MethodClientPCICEStat = "client:request:pc_ice_stat"
MethodClientTrace = "client:trace"
MethodClientConfSpeakersState = "client:request:conf_speakers_state_zip"
MethodServerConfSpeakersState = "server:response:conf_speakers_state_zip"
MethodClientGetVideoFromUser = "client:request:get_video_from_user"
MethodClientStopVideoFromUser = "client:request:stop_video_from_user"
MethodServerGetVideoFromUser = "server:response:get_video_from_user"
MethodServerStopVideoFromUser = "server:response:stop_video_from_user"
MethodClientCamStateChange = "client:request:cam_state_change"
MethodClientMicStateChange = "client:request:mic_state_change"
MethodClientScreenSharingSwitchOn = "client:request:screensharing_switch_on"
MethodClientScreenSharingSwitchOff = "client:request:screensharing_switch_off"
MethodClientGetScreenSharingFromUser = "client:request:get_screensharing_from_user"
MethodClientStopScreenSharingFromUser = "client:request:stop_screensharing_from_user"
MethodClientScreensharingQualityChange = "client:request:screensharing_quality_change"
MethodServerGetScreenSharingFromUser = "server:response:get_screensharing_from_user"
MethodClientClientStatZip = "client:request:client_stat_zip"
MethodClientKickOne = "client:request:kick_one"
MethodServerKickOneResponse = "server:response:kick_one"
MethodServerYouKicked = "server:you_kicked"
MethodServerYourCamStateChanged = "server:response:your_cam_state_changed"
MethodServerYourMicStateChanged = "server:response:your_mic_state_changed"
MethodServerSpeakerCamStateChanged = "server:speaker_cam_state_changed"
MethodServerSpeakerMicStateChanged = "server:speaker_mic_state_changed"
ProductVersion = "6.14.0"
SubscriptionVersion = "2.0"
)
type Frame struct {
JSONRPC string `json:"jsonrpc"`
Method string `json:"method,omitempty"`
Params json.RawMessage `json:"params,omitempty"`
Result json.RawMessage `json:"result,omitempty"`
Error *RPCError `json:"error,omitempty"`
ID json.RawMessage `json:"id,omitempty"`
}
type RPCError struct {
Code int `json:"code"`
Message string `json:"message"`
}
type TransceiverDesc struct {
TransceiverID string `json:"transceiver_id"`
SessionID string `json:"session_id"`
Direction string `json:"direction"`
Ctype string `json:"ctype"`
}
type DataChannelDesc struct {
ID int `json:"id"`
Label string `json:"label"`
}
type SDPEnvelope struct {
Type string `json:"type"`
SDP string `json:"sdp"`
}
type SDPOfferParams struct {
MicState bool `json:"mic_state"`
CamState bool `json:"cam_state"`
NoiseSuppressionState bool `json:"noise_suppression_state"`
VideoQuality *string `json:"video_quality"`
ScreenSharingQuality string `json:"screen_sharing_quality"`
Datachannels []DataChannelDesc `json:"datachannels"`
Transceivers []TransceiverDesc `json:"transceivers"`
Offer string `json:"offer"`
}
type SDPAnswerParams struct {
Answer string `json:"answer"`
Transceivers []TransceiverDesc `json:"transceivers"`
}
type ICEServerEntry struct {
URLs []string `json:"urls"`
Username string `json:"username"`
Credential string `json:"credential"`
}
type YouJoinedParams struct {
IcePolicy string `json:"ice_policy"`
IceServers []ICEServerEntry `json:"ice_servers"`
Event json.RawMessage `json:"event"`
EventParams json.RawMessage `json:"event_params"`
PreferredCodecs json.RawMessage `json:"preferred_codecs"`
}
type SpeakerJoinedParams struct {
SessionID string `json:"session_id"`
UserID string `json:"user_id,omitempty"`
Name string `json:"name,omitempty"`
CamState bool `json:"cam_state"`
MicState bool `json:"mic_state"`
Extra json.RawMessage `json:"-"`
}
type SpeakerCamStateChangedParams struct {
SessionID string `json:"session_id"`
CamState bool `json:"cam_state"`
}
type SpeakerMicStateChangedParams struct {
SessionID string `json:"session_id"`
MicState bool `json:"mic_state"`
}
type SpeakerEntry struct {
SessionID string `json:"session_id"`
UserID string `json:"user_id"`
Name string `json:"name"`
MicState bool `json:"mic_state"`
CamState bool `json:"cam_state"`
Role string `json:"role"`
WebinarRole string `json:"webinar_role"`
IsGuest bool `json:"is_guest"`
}
type ConfSpeakersStateResponse struct {
SpeakersCount int `json:"speakers_count"`
WebinarSpeakersCount int `json:"webinar_speakers_count"`
Speakers []SpeakerEntry `json:"speakers"`
}
type ConfSpeakersStateRequest struct {
SessionIDs []string `json:"session_ids"`
TileParams ConfSpeakersTileParams `json:"tile_params"`
InputVideoQuality string `json:"input_video_quality"`
ScreenParams ConfSpeakersScreenParams `json:"screen_params"`
}
type ConfSpeakersTileParams struct {
Mode string `json:"mode"`
MosaicParams ConfSpeakersMosaicParams `json:"mosaic_params"`
IsModeBlocked bool `json:"is_mode_blocked"`
}
type ConfSpeakersMosaicParams struct {
MaxTilesCount int `json:"max_tiles_count"`
}
type ConfSpeakersScreenParams struct {
Height int `json:"height"`
Width int `json:"width"`
}
type SpeakerDisconnectedParams struct {
SessionID string `json:"session_id"`
}
type ICECandidateJSON struct {
Candidate string `json:"candidate"`
SDPMid *string `json:"sdpMid"`
SDPMLineIndex *uint16 `json:"sdpMLineIndex"`
UsernameFragment string `json:"usernameFragment,omitempty"`
}
type GetVideoFromUserRequest struct {
SessionID string `json:"session_id"`
TransceiverID string `json:"transceiver_id"`
UserID string `json:"user_id"`
Username string `json:"username"`
}
type GetVideoFromUserResponse struct {
SessionID string `json:"session_id"`
TransceiverID string `json:"transceiver_id"`
}
type GetScreenSharingFromUserRequest struct {
SessionID string `json:"session_id"`
TransceiverID string `json:"transceiver_id"`
UserID string `json:"user_id"`
}
type GetScreenSharingFromUserResponse struct {
SessionID string `json:"session_id"`
TransceiverID string `json:"transceiver_id"`
}
type ClientStatVideoIn struct {
BytesReceived int64 `json:"bytes_received"`
Codec string `json:"codec"`
IsEnabled bool `json:"is_enabled"`
JitterBufferDelay float64 `json:"jitter_buffer_delay"`
JitterBufferEmittedCount int `json:"jitter_buffer_emitted_count"`
Jitter float64 `json:"jitter"`
Mid int `json:"mid"`
PacketsLost int `json:"packets_lost"`
PacketsReceived int `json:"packets_received"`
Framerate int `json:"framerate"`
FreezeCount int `json:"freeze_count"`
Resolution ClientStatResolution `json:"resolution"`
Rid string `json:"rid"`
TotalFreezesDuration int `json:"total_freezes_duration"`
SessionID string `json:"session_id"`
}
type ClientStatVideoOut struct {
Mid int `json:"mid"`
BytesSent int64 `json:"bytes_sent"`
Codec string `json:"codec"`
IsEnabled bool `json:"is_enabled"`
PacketsSent int `json:"packets_sent"`
RemoteStats ClientStatRemoteStats `json:"remote_stats"`
TargetBitrate int `json:"target_bitrate"`
Framerate int `json:"framerate"`
FreezeCount int `json:"freeze_count"`
Resolution ClientStatResolution `json:"resolution"`
Rid string `json:"rid"`
TotalFreezesDuration int `json:"total_freezes_duration"`
SessionID string `json:"session_id"`
ScalabilityMode string `json:"scalability_mode"`
}
type ClientStatResolution struct {
Height int `json:"height"`
Width int `json:"width"`
}
type ClientStatRemoteStats struct {
Jitter float64 `json:"jitter"`
FractionPacketsLost float64 `json:"fraction_packets_lost"`
PacketsLost int `json:"packets_lost"`
RTT float64 `json:"rtt"`
}
type ClientStatAudioIn struct {
BytesReceived int64 `json:"bytes_received"`
Codec string `json:"codec"`
IsEnabled bool `json:"is_enabled"`
JitterBufferDelay float64 `json:"jitter_buffer_delay"`
JitterBufferEmittedCount int `json:"jitter_buffer_emitted_count"`
Jitter float64 `json:"jitter"`
Mid int `json:"mid"`
PacketsLost int `json:"packets_lost"`
PacketsReceived int `json:"packets_received"`
}
type ClientStatConnection struct {
BytesReceived int64 `json:"bytes_received"`
BytesSent int64 `json:"bytes_sent"`
CurrentRTT float64 `json:"current_rtt"`
}
type ClientStatReport struct {
ReportTimeUnixMS int64 `json:"report_time_unix_ms"`
Connection ClientStatConnection `json:"connection"`
Audio struct {
In ClientStatAudioIn `json:"in"`
} `json:"audio"`
Video struct {
In []ClientStatVideoIn `json:"in"`
OutV2 []ClientStatVideoOut `json:"out_v2"`
Out ClientStatVideoOut `json:"out"`
} `json:"video"`
Screensharing struct{} `json:"screensharing"`
}
type SignalingDialOptions struct {
UserAgent string
Origin string
Logger logger.ContextLogger
Dialer N.Dialer
}
type SignalingClient struct {
conn *websocket.Conn
writeMu sync.Mutex
closed atomic.Bool
logger logger.ContextLogger
sessionID string
eventID string
OnYouJoined func(YouJoinedParams)
OnSubscribeResponse func()
OnSDPAnswer func(answerSDP string, transceivers []TransceiverDesc)
OnSpeakerJoined func(SpeakerJoinedParams)
OnSpeakerDisconnected func(SpeakerDisconnectedParams)
OnConfSpeakersState func(ConfSpeakersStateResponse)
OnSpeakerCamStateChanged func(SpeakerCamStateChangedParams)
OnSpeakerMicStateChanged func(SpeakerMicStateChangedParams)
OnGetVideoFromUserResponse func(resp GetVideoFromUserResponse, errCode int, errMessage string)
OnGetScreenSharingFromUserResponse func(resp GetScreenSharingFromUserResponse, errCode int, errMessage string)
OnHeartbeat func()
OnUnknown func(method string, params json.RawMessage)
OnDataChannelMessage func(method string, params json.RawMessage)
}
func DialSignaling(wssURL string, opts SignalingDialOptions) (*SignalingClient, error) {
if !strings.Contains(wssURL, "socket_version=") {
joiner := "&"
if !strings.Contains(wssURL, "?") {
joiner = "?"
}
wssURL = wssURL + joiner + "socket_version=2.0"
}
dialer := websocket.Dialer{HandshakeTimeout: 10 * time.Second}
if opts.Dialer != nil {
dialer.NetDialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
return opts.Dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
}
}
headers := http.Header{}
if opts.UserAgent != "" {
headers.Set("User-Agent", opts.UserAgent)
}
if opts.Origin != "" {
headers.Set("Origin", opts.Origin)
} else {
headers.Set("Origin", Origin)
}
log := opts.Logger
if log == nil {
log = logger.NOP()
}
conn, resp, err := dialer.Dial(wssURL, headers)
if err != nil {
status := 0
if resp != nil {
status = resp.StatusCode
}
return nil, fmt.Errorf("ws dial: %w status=%d url=%s", err, status, wssURL)
}
if resp != nil {
log.Debug(fmt.Sprintf("dion: ws dial status=%d", resp.StatusCode))
}
return &SignalingClient{conn: conn, logger: log}, nil
}
func (c *SignalingClient) Close() error {
if !c.closed.CompareAndSwap(false, true) {
return nil
}
common.CloseWS(c.conn)
return nil
}
func (c *SignalingClient) WaitConnected(timeout time.Duration) error {
c.conn.SetReadDeadline(time.Now().Add(timeout))
_, raw, err := c.conn.ReadMessage()
if err != nil {
return fmt.Errorf("read connected: %w", err)
}
var frame Frame
if err := json.Unmarshal(raw, &frame); err != nil {
return fmt.Errorf("decode connected: %w", err)
}
if frame.Method != MethodServerConnected {
return fmt.Errorf("expected %s, got %s", MethodServerConnected, frame.Method)
}
c.logger.Debug("dion: signaling connected")
return nil
}
func (c *SignalingClient) Subscribe(eventID, sessionID string) error {
c.eventID = eventID
c.sessionID = sessionID
return c.sendFrame(MethodClientSubscribeConference, map[string]any{
"event_id": eventID,
"conf_user_session_id": sessionID,
"main_user_session_id": nil,
"product_version": ProductVersion,
"subscription_version": SubscriptionVersion,
})
}
func (c *SignalingClient) SendTrace(deviceInfo map[string]any) error {
data, err := json.Marshal(deviceInfo)
if err != nil {
return fmt.Errorf("marshal trace: %w", err)
}
return c.sendFrame(MethodClientTrace, map[string]any{"data": string(data)})
}
func (c *SignalingClient) SendSDPOffer(params SDPOfferParams) error {
return c.sendFrame(MethodClientSDPOffer, params)
}
func (c *SignalingClient) SendConfSpeakersState(request ConfSpeakersStateRequest) error {
encoded, err := ZipEncode(request)
if err != nil {
return fmt.Errorf("zip conf_speakers_state: %w", err)
}
return c.sendFrame(MethodClientConfSpeakersState, encoded)
}
func (c *SignalingClient) SendGetVideoFromUser(request GetVideoFromUserRequest) error {
return c.sendFrame(MethodClientGetVideoFromUser, request)
}
func (c *SignalingClient) SendStopVideoFromUser(request GetVideoFromUserRequest) error {
return c.sendFrame(MethodClientStopVideoFromUser, request)
}
func (c *SignalingClient) SendCamStateChange(state bool) error {
return c.sendFrame(MethodClientCamStateChange, map[string]any{"state": state})
}
func (c *SignalingClient) SendMicStateChange(state bool) error {
return c.sendFrame(MethodClientMicStateChange, map[string]any{"state": state})
}
func (c *SignalingClient) SendScreenSharingSwitchOn() error {
return c.sendFrame(MethodClientScreenSharingSwitchOn, map[string]any{})
}
func (c *SignalingClient) SendScreenSharingSwitchOff() error {
return c.sendFrame(MethodClientScreenSharingSwitchOff, map[string]any{})
}
func (c *SignalingClient) SendGetScreenSharingFromUser(request GetScreenSharingFromUserRequest) error {
return c.sendFrame(MethodClientGetScreenSharingFromUser, request)
}
func (c *SignalingClient) SendScreensharingQualityChange(quality string) error {
return c.sendFrame(MethodClientScreensharingQualityChange, map[string]any{"quality": quality})
}
func (c *SignalingClient) SendKickOne(sessionID string) error {
return c.sendFrame(MethodClientKickOne, map[string]any{"session_id": sessionID})
}
func (c *SignalingClient) SendPCIceStat() error {
return c.sendFrame(MethodClientPCICEStat, map[string]any{"device": "web"})
}
func (c *SignalingClient) SendClientStatZip(report ClientStatReport) error {
encoded, err := ZipEncode(report)
if err != nil {
return fmt.Errorf("zip client_stat: %w", err)
}
return c.sendFrame(MethodClientClientStatZip, encoded)
}
func (c *SignalingClient) SendICECandidates(candidates []ICECandidateJSON) error {
encoded := make([]string, 0, len(candidates))
for _, candidate := range candidates {
raw, err := EncodeICECandidate(candidate)
if err != nil {
return fmt.Errorf("encode candidate: %w", err)
}
encoded = append(encoded, raw)
}
zipped, err := ZipEncode(map[string]any{"candidates": encoded})
if err != nil {
return fmt.Errorf("zip candidates: %w", err)
}
return c.sendFrame(MethodClientSendICECandidates, zipped)
}
func (c *SignalingClient) ReadLoop() error {
for {
if c.closed.Load() {
return nil
}
c.conn.SetReadDeadline(time.Now().Add(60 * time.Second))
_, raw, err := c.conn.ReadMessage()
if err != nil {
if c.closed.Load() {
return nil
}
return fmt.Errorf("ws read: %w", err)
}
var frame Frame
if err := json.Unmarshal(raw, &frame); err != nil {
c.logger.Debug(fmt.Sprintf("dion: drop non-json frame: %v", err))
continue
}
if frame.Error != nil {
c.logger.Debug(fmt.Sprintf("dion: <- %s ERROR code=%d message=%q", frame.Method, frame.Error.Code, frame.Error.Message))
}
c.dispatch(frame)
}
}
func EncodeICECandidate(candidate ICECandidateJSON) (string, error) {
plain, err := json.Marshal(candidate)
if err != nil {
return "", err
}
return base64.StdEncoding.EncodeToString(plain), nil
}
func DecodeICECandidate(encoded string) (ICECandidateJSON, error) {
raw, err := base64.StdEncoding.DecodeString(encoded)
if err != nil {
return ICECandidateJSON{}, fmt.Errorf("base64: %w", err)
}
var out ICECandidateJSON
if err := json.Unmarshal(raw, &out); err != nil {
return ICECandidateJSON{}, fmt.Errorf("unmarshal: %w", err)
}
return out, nil
}
func BuildSDPOfferEnvelope(offerSDP string) (string, error) {
return ZipEncode(SDPEnvelope{Type: "offer", SDP: offerSDP})
}
func DecodeSDPAnswerInner(answerZipped string) (string, error) {
var inner SDPEnvelope
if err := ZipDecode(answerZipped, &inner); err != nil {
return "", err
}
return inner.SDP, nil
}
func DefaultConfSpeakersStateRequest() ConfSpeakersStateRequest {
return ConfSpeakersStateRequest{
SessionIDs: []string{},
TileParams: ConfSpeakersTileParams{
Mode: "mosaic",
MosaicParams: ConfSpeakersMosaicParams{MaxTilesCount: 9},
IsModeBlocked: false,
},
InputVideoQuality: "auto",
ScreenParams: ConfSpeakersScreenParams{Height: 720, Width: 1280},
}
}
func (c *SignalingClient) sendFrame(method string, params any) error {
c.writeMu.Lock()
defer c.writeMu.Unlock()
if c.closed.Load() {
return fmt.Errorf("signaling closed")
}
payload := map[string]any{
"jsonrpc": "2.0",
"method": method,
"params": params,
}
raw, err := json.Marshal(payload)
if err != nil {
return fmt.Errorf("marshal frame: %w", err)
}
return c.conn.WriteMessage(websocket.TextMessage, raw)
}
func (c *SignalingClient) dispatch(frame Frame) {
switch frame.Method {
case MethodServerConnected:
c.logger.Debug("dion: late server:notify:main:connected")
case MethodServerYouJoined:
var youJoined YouJoinedParams
if err := json.Unmarshal(frame.Params, &youJoined); err != nil {
c.logger.Debug(fmt.Sprintf("dion: decode you_joined: %v", err))
return
}
if c.OnYouJoined != nil {
c.OnYouJoined(youJoined)
}
case MethodServerSubscribeResponse:
if c.OnSubscribeResponse != nil {
c.OnSubscribeResponse()
}
case MethodServerSDPAnswer:
var answer SDPAnswerParams
if err := json.Unmarshal(frame.Params, &answer); err != nil {
c.logger.Debug(fmt.Sprintf("dion: decode sdp_answer: %v", err))
return
}
var inner SDPEnvelope
if err := ZipDecode(answer.Answer, &inner); err != nil {
c.logger.Debug(fmt.Sprintf("dion: decode sdp_answer envelope: %v", err))
return
}
if c.OnSDPAnswer != nil {
c.OnSDPAnswer(inner.SDP, answer.Transceivers)
}
case MethodServerSpeakerJoined:
var joined SpeakerJoinedParams
if err := json.Unmarshal(frame.Params, &joined); err != nil {
c.logger.Debug(fmt.Sprintf("dion: decode speaker_joined: %v", err))
return
}
joined.Extra = frame.Params
if c.OnSpeakerJoined != nil {
c.OnSpeakerJoined(joined)
}
case MethodServerSpeakerDisconnected:
var left SpeakerDisconnectedParams
if err := json.Unmarshal(frame.Params, &left); err != nil {
c.logger.Debug(fmt.Sprintf("dion: decode speaker_disconnected: %v", err))
return
}
if c.OnSpeakerDisconnected != nil {
c.OnSpeakerDisconnected(left)
}
case MethodServerSpeakerCamStateChanged:
var changed SpeakerCamStateChangedParams
if err := json.Unmarshal(frame.Params, &changed); err != nil {
c.logger.Debug(fmt.Sprintf("dion: decode speaker_cam_state_changed: %v", err))
return
}
if c.OnSpeakerCamStateChanged != nil {
c.OnSpeakerCamStateChanged(changed)
}
case MethodServerSpeakerMicStateChanged:
var changed SpeakerMicStateChangedParams
if err := json.Unmarshal(frame.Params, &changed); err != nil {
c.logger.Debug(fmt.Sprintf("dion: decode speaker_mic_state_changed: %v", err))
return
}
if c.OnSpeakerMicStateChanged != nil {
c.OnSpeakerMicStateChanged(changed)
}
case MethodServerConfSpeakersState:
var encoded string
if err := json.Unmarshal(frame.Params, &encoded); err != nil {
c.logger.Debug(fmt.Sprintf("dion: decode conf_speakers_state envelope: %v", err))
return
}
var response ConfSpeakersStateResponse
if err := ZipDecode(encoded, &response); err != nil {
c.logger.Debug(fmt.Sprintf("dion: decode conf_speakers_state body: %v", err))
return
}
if c.OnConfSpeakersState != nil {
c.OnConfSpeakersState(response)
}
case MethodServerHeartbeat:
if c.OnHeartbeat != nil {
c.OnHeartbeat()
}
case MethodServerGetVideoFromUser:
var resp GetVideoFromUserResponse
_ = json.Unmarshal(frame.Params, &resp)
errCode := 0
errMsg := ""
if frame.Error != nil {
errCode = frame.Error.Code
errMsg = frame.Error.Message
}
if c.OnGetVideoFromUserResponse != nil {
c.OnGetVideoFromUserResponse(resp, errCode, errMsg)
}
case MethodServerGetScreenSharingFromUser:
var resp GetScreenSharingFromUserResponse
_ = json.Unmarshal(frame.Params, &resp)
errCode := 0
errMsg := ""
if frame.Error != nil {
errCode = frame.Error.Code
errMsg = frame.Error.Message
}
if c.OnGetScreenSharingFromUserResponse != nil {
c.OnGetScreenSharingFromUserResponse(resp, errCode, errMsg)
}
default:
if c.OnUnknown != nil {
c.OnUnknown(frame.Method, frame.Params)
}
}
}

View File

@@ -0,0 +1,49 @@
package dion
import (
"bytes"
"compress/gzip"
"encoding/base64"
"encoding/json"
"fmt"
"io"
)
func ZipEncode(value any) (string, error) {
plain, err := json.Marshal(value)
if err != nil {
return "", fmt.Errorf("marshal: %w", err)
}
var compressed bytes.Buffer
gz := gzip.NewWriter(&compressed)
if _, err := gz.Write(plain); err != nil {
return "", fmt.Errorf("gzip write: %w", err)
}
if err := gz.Close(); err != nil {
return "", fmt.Errorf("gzip close: %w", err)
}
return base64.StdEncoding.EncodeToString(compressed.Bytes()), nil
}
func ZipDecode(encoded string, out any) error {
raw, err := base64.StdEncoding.DecodeString(encoded)
if err != nil {
return fmt.Errorf("base64: %w", err)
}
reader, err := gzip.NewReader(bytes.NewReader(raw))
if err != nil {
return fmt.Errorf("gzip reader: %w", err)
}
defer reader.Close()
plain, err := io.ReadAll(reader)
if err != nil {
return fmt.Errorf("gzip read: %w", err)
}
if out == nil {
return nil
}
if err := json.Unmarshal(plain, out); err != nil {
return fmt.Errorf("unmarshal: %w", err)
}
return nil
}

View File

@@ -0,0 +1,472 @@
package livekit
import (
"context"
"encoding/json"
"fmt"
"net"
"net/http"
"net/netip"
"net/url"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/gorilla/websocket"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/dialer"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
const (
ProtocolVersion = "15"
SDKName = "js"
SDKVersion = "2.7.0"
PingPeriod = 5 * time.Second
TargetPublisher = signalTargetPublisher
TargetSubscriber = signalTargetSubscriber
TrackTypeAudio = trackTypeAudio
TrackTypeVideo = trackTypeVideo
TrackTypeData = trackTypeData
TrackSourceCamera = trackSourceCamera
TrackSourceScreenShare = trackSourceScreenShare
)
type ICEServer = iceServer
type JoinResponse = joinResponse
type Config struct {
ServerURL string
Token string
Origin string
UserAgent string
Logger logger.ContextLogger
SettingEngine *webrtc.SettingEngine
NetDialContext func(ctx context.Context, network, addr string) (net.Conn, error)
DNSRouter adapter.DNSRouter
Dialer N.Dialer
}
type Client struct {
logger logger.ContextLogger
wsURL string
token string
origin string
ua string
settingEngine *webrtc.SettingEngine
netDialContext func(ctx context.Context, network, addr string) (net.Conn, error)
dnsRouter adapter.DNSRouter
dialer N.Dialer
ws *websocket.Conn
wsMu sync.Mutex
join JoinResponse
pubPC *webrtc.PeerConnection
subPC *webrtc.PeerConnection
pubMu sync.Mutex
subMu sync.Mutex
pubRemoteSet bool
subRemoteSet bool
closed atomic.Bool
OnReady func()
OnTrack func(*webrtc.TrackRemote, *webrtc.RTPReceiver)
OnDataChannel func(*webrtc.DataChannel)
OnPubConnected func()
OnParticipantUpdate func([]ParticipantInfo)
OnRemoteCandidate func(target int, candidate webrtc.ICECandidateInit)
OnRemoteSDP func(target int, sdpType, sdp string)
}
func NewClient(cfg Config) *Client {
return &Client{
logger: cfg.Logger,
wsURL: cfg.ServerURL,
token: cfg.Token,
origin: cfg.Origin,
ua: cfg.UserAgent,
settingEngine: cfg.SettingEngine,
netDialContext: cfg.NetDialContext,
dnsRouter: cfg.DNSRouter,
dialer: cfg.Dialer,
}
}
func (c *Client) Join() JoinResponse { return c.join }
func (c *Client) PubPC() *webrtc.PeerConnection { return c.pubPC }
func (c *Client) SubPC() *webrtc.PeerConnection { return c.subPC }
func (c *Client) Connect() error {
u, err := url.Parse(c.wsURL)
if err != nil {
return fmt.Errorf("parse url: %w", err)
}
u.Path = "/rtc"
q := u.Query()
q.Set("access_token", c.token)
q.Set("protocol", ProtocolVersion)
q.Set("sdk", SDKName)
q.Set("version", SDKVersion)
q.Set("auto_subscribe", "1")
q.Set("adaptive_stream", "true")
u.RawQuery = q.Encode()
headers := http.Header{}
if c.ua != "" {
headers.Set("User-Agent", c.ua)
}
if c.origin != "" {
headers.Set("Origin", c.origin)
}
dialer := *websocket.DefaultDialer
if c.netDialContext != nil {
dialer.NetDialContext = c.netDialContext
}
conn, resp, err := dialer.Dial(u.String(), headers)
if err != nil {
if resp != nil {
return fmt.Errorf("ws dial: %w (status %d)", err, resp.StatusCode)
}
return fmt.Errorf("ws dial: %w", err)
}
c.ws = conn
c.logger.Info("[lk] signaling connected")
return nil
}
func (c *Client) SendOffer(sdp string) error {
return c.sendSignal(encSignalRequestOffer(sessionDescription{Type: "offer", SDP: sdp}))
}
func (c *Client) SendAnswer(sdp string) error {
return c.sendSignal(encSignalRequestAnswer(sessionDescription{Type: "answer", SDP: sdp}))
}
func (c *Client) SendTrickle(candidate webrtc.ICECandidateInit, target int) error {
js, _ := json.Marshal(candidate)
return c.sendSignal(encSignalRequestTrickle(trickleMsg{
CandidateInit: string(js),
Target: target,
}))
}
func (c *Client) SendAddTrack(cid, name string, trackType, source int, width, height uint32) error {
return c.sendSignal(encSignalRequestAddTrack(cid, name, trackType, source, width, height))
}
func (c *Client) SendLeave() error { return c.sendSignal(encSignalRequestLeave()) }
func (c *Client) SendPing() error {
return c.sendSignal(encSignalRequestPing(time.Now().UnixMilli()))
}
func (c *Client) Close() {
if !c.closed.CompareAndSwap(false, true) {
return
}
c.wsMu.Lock()
ws := c.ws
c.wsMu.Unlock()
common.CloseWS(ws)
if c.pubPC != nil {
_ = c.pubPC.Close()
}
if c.subPC != nil {
_ = c.subPC.Close()
}
}
func (c *Client) ReadLoop() error {
defer c.Close()
for {
mt, data, err := c.ws.ReadMessage()
if err != nil {
return err
}
if mt != websocket.BinaryMessage {
continue
}
c.handleSignal(data)
}
}
func (c *Client) PingLoop() {
period := PingPeriod
if c.join.PingIntervalSec > 0 {
period = time.Duration(c.join.PingIntervalSec) * time.Second
}
t := time.NewTicker(period)
defer t.Stop()
var sentN int
for range t.C {
if c.closed.Load() {
return
}
if err := c.SendPing(); err != nil {
c.logger.Warn(fmt.Sprintf("[lk] ping send failed: %v", err))
return
}
sentN++
if sentN <= 3 || sentN%12 == 0 {
c.logger.Debug(fmt.Sprintf("[lk] ping #%d sent", sentN))
}
}
}
func (c *Client) sendSignal(payload []byte) error {
c.wsMu.Lock()
defer c.wsMu.Unlock()
if c.ws == nil {
return fmt.Errorf("ws not connected")
}
return c.ws.WriteMessage(websocket.BinaryMessage, payload)
}
func (c *Client) iceServersAsWebRTC() []webrtc.ICEServer {
out := make([]webrtc.ICEServer, 0, len(c.join.ICEServers))
resolved := make(map[string]string)
for _, s := range c.join.ICEServers {
urls := make([]string, len(s.URLs))
copy(urls, s.URLs)
for k, u := range urls {
host := common.ExtractICEHost(u)
if host == "" || net.ParseIP(host) != nil {
continue
}
ip, ok := resolved[host]
if !ok {
rd, hasRD := c.dialer.(dialer.ResolveDialer)
if c.dnsRouter == nil || !hasRD {
continue
}
var addrs []netip.Addr
var err error
addrs, err = c.dnsRouter.Lookup(context.Background(), host, rd.QueryOptions())
if err != nil {
c.logger.Warn(fmt.Sprintf("[lk] resolve ICE host %s failed: %v", host, err))
continue
}
resolved[host] = addrs[0].String()
c.logger.Debug(fmt.Sprintf("[lk] resolved ICE host %s -> %s", host, addrs[0]))
}
urls[k] = strings.Replace(u, host, ip, 1)
}
ice := webrtc.ICEServer{URLs: urls}
if s.Username != "" {
ice.Username = s.Username
ice.Credential = s.Credential
}
out = append(out, ice)
}
return out
}
func (c *Client) buildPeerConnections() error {
cfg := webrtc.Configuration{ICEServers: c.iceServersAsWebRTC()}
se := webrtc.SettingEngine{}
if c.settingEngine != nil {
se = *c.settingEngine
}
se.DetachDataChannels()
api := webrtc.NewAPI(webrtc.WithSettingEngine(se))
pubPC, err := api.NewPeerConnection(cfg)
if err != nil {
return fmt.Errorf("create pub pc: %w", err)
}
subPC, err := api.NewPeerConnection(cfg)
if err != nil {
_ = pubPC.Close()
return fmt.Errorf("create sub pc: %w", err)
}
c.pubPC = pubPC
c.subPC = subPC
pubPC.OnICECandidate(func(cand *webrtc.ICECandidate) {
if cand == nil {
c.logger.Debug("[lk] pub ICE gathering complete")
return
}
c.logger.Debug(fmt.Sprintf("[lk] pub local cand: %s", cand.String()))
_ = c.SendTrickle(cand.ToJSON(), TargetPublisher)
})
subPC.OnICECandidate(func(cand *webrtc.ICECandidate) {
if cand == nil {
c.logger.Debug("[lk] sub ICE gathering complete")
return
}
c.logger.Debug(fmt.Sprintf("[lk] sub local cand: %s", cand.String()))
_ = c.SendTrickle(cand.ToJSON(), TargetSubscriber)
})
pubPC.OnConnectionStateChange(func(state webrtc.PeerConnectionState) {
c.logger.Debug(fmt.Sprintf("[lk] pub PC state: %s", state.String()))
if state == webrtc.PeerConnectionStateConnected && c.OnPubConnected != nil {
c.OnPubConnected()
}
})
subPC.OnConnectionStateChange(func(state webrtc.PeerConnectionState) {
c.logger.Debug(fmt.Sprintf("[lk] sub PC state: %s", state.String()))
})
pubPC.OnICEConnectionStateChange(func(state webrtc.ICEConnectionState) {
c.logger.Debug(fmt.Sprintf("[lk] pub ICE state: %s", state.String()))
})
subPC.OnICEConnectionStateChange(func(state webrtc.ICEConnectionState) {
c.logger.Debug(fmt.Sprintf("[lk] sub ICE state: %s", state.String()))
})
subPC.OnTrack(func(track *webrtc.TrackRemote, receiver *webrtc.RTPReceiver) {
c.logger.Debug(fmt.Sprintf("[lk] sub remote track: %s", track.Codec().MimeType))
if c.OnTrack != nil {
c.OnTrack(track, receiver)
}
})
subPC.OnDataChannel(func(dc *webrtc.DataChannel) {
c.logger.Debug(fmt.Sprintf("[lk] sub data channel: %s", dc.Label()))
if c.OnDataChannel != nil {
c.OnDataChannel(dc)
}
})
c.logger.Debug(fmt.Sprintf("[lk] PCs created (%d ICE servers)", len(c.join.ICEServers)))
for i, s := range c.join.ICEServers {
c.logger.Debug(fmt.Sprintf("[lk] iceServer[%d]: urls=%v hasCred=%v", i, s.URLs, s.Username != ""))
}
return nil
}
func (c *Client) handleSignal(data []byte) {
sr, err := decSignalResponse(data)
if err != nil {
c.logger.Warn(fmt.Sprintf("[lk] decode signal: %v", err))
return
}
switch sr.Kind {
case signalRespJoin:
if sr.Join != nil {
c.join = *sr.Join
c.logger.Info(fmt.Sprintf("[lk] join: room=%s participant=%s subscriberPrimary=%v iceServers=%d pingTimeout=%ds pingInterval=%ds",
c.join.RoomName, c.join.ParticipantID, c.join.SubscriberPrimary, len(c.join.ICEServers),
c.join.PingTimeoutSec, c.join.PingIntervalSec))
if err := c.buildPeerConnections(); err != nil {
c.logger.Error(fmt.Sprintf("[lk] %v", err))
return
}
if c.OnReady != nil {
c.OnReady()
}
}
case signalRespAnswer:
c.logger.Debug(fmt.Sprintf("[lk] <- pub answer (%d bytes)", len(sr.SDP.SDP)))
if sr.SDP != nil {
c.applyPubAnswer(sr.SDP.SDP)
}
case signalRespOffer:
c.logger.Debug(fmt.Sprintf("[lk] <- sub offer (%d bytes)", len(sr.SDP.SDP)))
if sr.SDP != nil {
c.applySubOfferAndAnswer(sr.SDP.SDP)
}
case signalRespTrickle:
if sr.Trickle != nil {
c.logger.Debug(fmt.Sprintf("[lk] <- trickle target=%d", sr.Trickle.Target))
c.applyRemoteTrickle(*sr.Trickle)
}
case signalRespRefreshToken:
if sr.Token != "" {
c.token = sr.Token
c.logger.Debug("[lk] token refreshed")
}
case signalRespLeave:
if sr.Leave != nil {
c.logger.Debug(fmt.Sprintf("[lk] ignored leave reason=%s action=%s",
DisconnectReasonName(sr.Leave.Reason), LeaveActionName(sr.Leave.Action)))
} else {
c.logger.Debug("[lk] ignored leave")
}
case signalRespUpdate:
if c.OnParticipantUpdate != nil && len(sr.Participants) > 0 {
c.OnParticipantUpdate(sr.Participants)
}
default:
c.logger.Debug(fmt.Sprintf("[lk] <- signal kind=%d (%d bytes)", sr.Kind, len(data)))
}
}
func (c *Client) applyPubAnswer(sdp string) {
if c.OnRemoteSDP != nil {
c.OnRemoteSDP(TargetPublisher, "answer", sdp)
}
c.pubMu.Lock()
defer c.pubMu.Unlock()
if c.pubPC == nil {
return
}
if err := c.pubPC.SetRemoteDescription(webrtc.SessionDescription{Type: webrtc.SDPTypeAnswer, SDP: sdp}); err != nil {
c.logger.Warn(fmt.Sprintf("[lk] set pub remote answer: %v", err))
return
}
c.pubRemoteSet = true
}
func (c *Client) applySubOfferAndAnswer(sdp string) {
if c.OnRemoteSDP != nil {
c.OnRemoteSDP(TargetSubscriber, "offer", sdp)
}
c.subMu.Lock()
defer c.subMu.Unlock()
if c.subPC == nil {
return
}
if err := c.subPC.SetRemoteDescription(webrtc.SessionDescription{Type: webrtc.SDPTypeOffer, SDP: sdp}); err != nil {
c.logger.Warn(fmt.Sprintf("[lk] set sub remote offer: %v", err))
return
}
c.subRemoteSet = true
answer, err := c.subPC.CreateAnswer(nil)
if err != nil {
c.logger.Warn(fmt.Sprintf("[lk] create sub answer: %v", err))
return
}
if err := c.subPC.SetLocalDescription(answer); err != nil {
c.logger.Warn(fmt.Sprintf("[lk] set sub local answer: %v", err))
return
}
if err := c.SendAnswer(answer.SDP); err != nil {
c.logger.Warn(fmt.Sprintf("[lk] send answer: %v", err))
}
}
func (c *Client) applyRemoteTrickle(m trickleMsg) {
if m.CandidateInit == "" {
return
}
var ic webrtc.ICECandidateInit
if err := json.Unmarshal([]byte(m.CandidateInit), &ic); err != nil {
c.logger.Warn(fmt.Sprintf("[lk] decode trickle candidate: %v", err))
return
}
if c.OnRemoteCandidate != nil {
c.OnRemoteCandidate(m.Target, ic)
}
switch m.Target {
case TargetPublisher:
c.pubMu.Lock()
ready := c.pubRemoteSet
c.pubMu.Unlock()
if ready {
_ = c.pubPC.AddICECandidate(ic)
}
case TargetSubscriber:
c.subMu.Lock()
ready := c.subRemoteSet
c.subMu.Unlock()
if ready {
_ = c.subPC.AddICECandidate(ic)
}
}
}

View File

@@ -0,0 +1,708 @@
package livekit
import "fmt"
const (
signalReqOffer = 1
signalReqAnswer = 2
signalReqTrickle = 3
signalReqAddTrack = 4
signalReqLeave = 8
signalReqPingLegacy = 14
signalReqPingReq = 16
signalRespJoin = 1
signalRespAnswer = 2
signalRespOffer = 3
signalRespTrickle = 4
signalRespUpdate = 5
signalRespTrackPublished = 6
signalRespLeave = 8
signalRespRoomUpdate = 11
signalRespRefreshToken = 16
signalRespPongResp = 20
signalRespRequestResponse = 22
signalRespTrackSubscribed = 23
sdpFieldType = 1
sdpFieldSDP = 2
sdpFieldID = 3
trickleFieldCandidate = 1
trickleFieldTarget = 2
trickleFieldFinal = 3
addTrackFieldCID = 1
addTrackFieldName = 2
addTrackFieldType = 3
addTrackFieldWidth = 4
addTrackFieldHeight = 5
addTrackFieldSource = 8
addTrackFieldLayers = 9
videoLayerFieldQuality = 1
videoLayerFieldWidth = 2
videoLayerFieldHeight = 3
videoQualityHigh = 2
joinFieldRoom = 1
joinFieldParticipant = 2
joinFieldOtherParticipants = 3
joinFieldServerVersion = 4
joinFieldICEServers = 5
joinFieldSubscriberPrimary = 6
joinFieldServerRegion = 9
joinFieldPingTimeout = 10
joinFieldPingInterval = 11
iceServerFieldURLs = 1
iceServerFieldUsername = 2
iceServerFieldCredential = 3
pingFieldTimestamp = 1
pingFieldRTT = 2
dataPacketFieldKind = 1
dataPacketFieldUser = 2
userPacketFieldPayload = 2
DataPacketKindReliable = 0
DataPacketKindLossy = 1
leaveFieldCanReconnect = 1
leaveFieldReason = 2
leaveFieldAction = 3
roomFieldSID = 1
roomFieldName = 2
participantFieldSID = 1
participantFieldIdentity = 2
participantFieldState = 3
participantFieldName = 9
trackTypeAudio = 0
trackTypeVideo = 1
trackTypeData = 2
signalTargetPublisher = 0
signalTargetSubscriber = 1
trackSourceCamera = 1
trackSourceScreenShare = 3
)
const (
ParticipantStateJoining int32 = 0
ParticipantStateJoined int32 = 1
ParticipantStateActive int32 = 2
ParticipantStateDisconnected int32 = 3
)
var disconnectReasonNames = map[int]string{
0: "UNKNOWN",
1: "CLIENT_INITIATED",
2: "DUPLICATE_IDENTITY",
3: "SERVER_SHUTDOWN",
4: "PARTICIPANT_REMOVED",
5: "ROOM_DELETED",
6: "STATE_MISMATCH",
7: "JOIN_FAILURE",
8: "MIGRATION",
9: "SIGNAL_CLOSE",
10: "ROOM_CLOSED",
11: "USER_UNAVAILABLE",
12: "USER_REJECTED",
13: "SIP_TRUNK_FAILURE",
14: "CONNECTION_TIMEOUT",
15: "MEDIA_FAILURE",
16: "AGENT_ERROR",
}
var leaveActionNames = map[int]string{
0: "DISCONNECT",
1: "RESUME",
2: "RECONNECT",
}
type LeaveInfo struct {
Reason int
Action int
}
type sessionDescription struct {
Type string
SDP string
ID uint32
}
type trickleMsg struct {
CandidateInit string
Target int
Final bool
}
type iceServer struct {
URLs []string
Username string
Credential string
}
type joinResponse struct {
RoomSID string
RoomName string
ParticipantSID string
ParticipantID string
ServerVersion string
ServerRegion string
ICEServers []iceServer
SubscriberPrimary bool
PingTimeoutSec int32
PingIntervalSec int32
}
type signalResponse struct {
Kind int
Join *joinResponse
SDP *sessionDescription
Trickle *trickleMsg
Token string
PongTime int64
Leave *LeaveInfo
Participants []ParticipantInfo
}
type ParticipantInfo struct {
SID string
Identity string
State int32
Name string
}
func DisconnectReasonName(code int) string {
if name, ok := disconnectReasonNames[code]; ok {
return name
}
return fmt.Sprintf("CODE_%d", code)
}
func LeaveActionName(code int) string {
if name, ok := leaveActionNames[code]; ok {
return name
}
return fmt.Sprintf("CODE_%d", code)
}
func DecodeLeaveRequest(data []byte) LeaveInfo {
r := pbReader{buf: data}
var li LeaveInfo
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return li
}
switch {
case field == leaveFieldReason && wire == wireVarint:
v, _ := r.varint()
li.Reason = int(v)
case field == leaveFieldAction && wire == wireVarint:
v, _ := r.varint()
li.Action = int(v)
default:
if err := r.skipWire(wire); err != nil {
return li
}
}
}
return li
}
func EncodeDataPacketUser(payload []byte, kind int) []byte {
w := pbWriter{}
if kind != 0 {
w.int32(dataPacketFieldKind, int32(kind))
}
w.message(dataPacketFieldUser, encUserPacket(payload))
return w.buf
}
func DecodeDataPacketUser(data []byte) ([]byte, bool) {
r := pbReader{buf: data}
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return nil, false
}
if field == dataPacketFieldUser && wire == wireBytes {
inner, err := r.bytes()
if err != nil {
return nil, false
}
ur := pbReader{buf: inner}
for !ur.eof() {
ufield, uwire, uerr := ur.tag()
if uerr != nil {
return nil, false
}
if ufield == userPacketFieldPayload && uwire == wireBytes {
payload, perr := ur.bytes()
if perr != nil {
return nil, false
}
out := make([]byte, len(payload))
copy(out, payload)
return out, true
}
if err := ur.skipWire(uwire); err != nil {
return nil, false
}
}
return nil, false
}
if err := r.skipWire(wire); err != nil {
return nil, false
}
}
return nil, false
}
func DecodeParticipantInfo(data []byte) ParticipantInfo {
r := pbReader{buf: data}
var info ParticipantInfo
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return info
}
switch {
case field == participantFieldSID && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return info
}
info.SID = string(b)
case field == participantFieldIdentity && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return info
}
info.Identity = string(b)
case field == participantFieldState && wire == wireVarint:
v, err := r.varint()
if err != nil {
return info
}
info.State = int32(v)
case field == participantFieldName && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return info
}
info.Name = string(b)
default:
if err := r.skipWire(wire); err != nil {
return info
}
}
}
return info
}
func DecodeParticipantUpdate(data []byte) []ParticipantInfo {
r := pbReader{buf: data}
var out []ParticipantInfo
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return out
}
if field == 1 && wire == wireBytes {
b, err := r.bytes()
if err != nil {
return out
}
out = append(out, DecodeParticipantInfo(b))
} else {
if err := r.skipWire(wire); err != nil {
return out
}
}
}
return out
}
func encSessionDescription(sd sessionDescription) []byte {
w := pbWriter{}
if sd.Type != "" {
w.string(sdpFieldType, sd.Type)
}
if sd.SDP != "" {
w.string(sdpFieldSDP, sd.SDP)
}
if sd.ID != 0 {
w.uint32(sdpFieldID, sd.ID)
}
return w.buf
}
func encTrickle(m trickleMsg) []byte {
w := pbWriter{}
w.string(trickleFieldCandidate, m.CandidateInit)
w.int32(trickleFieldTarget, int32(m.Target))
if m.Final {
w.bool(trickleFieldFinal, true)
}
return w.buf
}
func encVideoLayer(quality, width, height uint32) []byte {
w := pbWriter{}
if quality != 0 {
w.uint32(videoLayerFieldQuality, quality)
}
if width != 0 {
w.uint32(videoLayerFieldWidth, width)
}
if height != 0 {
w.uint32(videoLayerFieldHeight, height)
}
return w.buf
}
func encAddTrack(cid, name string, trackType, source int, width, height uint32) []byte {
w := pbWriter{}
w.string(addTrackFieldCID, cid)
w.string(addTrackFieldName, name)
w.int32(addTrackFieldType, int32(trackType))
if width != 0 {
w.uint32(addTrackFieldWidth, width)
}
if height != 0 {
w.uint32(addTrackFieldHeight, height)
}
w.int32(addTrackFieldSource, int32(source))
if trackType == trackTypeVideo {
w.message(addTrackFieldLayers, encVideoLayer(videoQualityHigh, width, height))
}
return w.buf
}
func encPing(timestamp int64) []byte {
w := pbWriter{}
w.int64(pingFieldTimestamp, timestamp)
return w.buf
}
func encUserPacket(payload []byte) []byte {
w := pbWriter{}
w.bytes(userPacketFieldPayload, payload)
return w.buf
}
func encSignalRequestOffer(sd sessionDescription) []byte {
w := pbWriter{}
w.message(signalReqOffer, encSessionDescription(sd))
return w.buf
}
func encSignalRequestAnswer(sd sessionDescription) []byte {
w := pbWriter{}
w.message(signalReqAnswer, encSessionDescription(sd))
return w.buf
}
func encSignalRequestTrickle(m trickleMsg) []byte {
w := pbWriter{}
w.message(signalReqTrickle, encTrickle(m))
return w.buf
}
func encSignalRequestAddTrack(cid, name string, trackType, source int, width, height uint32) []byte {
w := pbWriter{}
w.message(signalReqAddTrack, encAddTrack(cid, name, trackType, source, width, height))
return w.buf
}
func encSignalRequestLeave() []byte {
w := pbWriter{}
w.message(signalReqLeave, []byte{})
return w.buf
}
func encSignalRequestPing(timestamp int64) []byte {
w := pbWriter{}
w.int64(signalReqPingLegacy, timestamp)
w.message(signalReqPingReq, encPing(timestamp))
return w.buf
}
func decSessionDescription(data []byte) (sessionDescription, error) {
r := pbReader{buf: data}
var sd sessionDescription
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return sd, err
}
switch {
case field == sdpFieldType && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return sd, err
}
sd.Type = string(b)
case field == sdpFieldSDP && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return sd, err
}
sd.SDP = string(b)
case field == sdpFieldID && wire == wireVarint:
v, err := r.varint()
if err != nil {
return sd, err
}
sd.ID = uint32(v)
default:
if err := r.skipWire(wire); err != nil {
return sd, err
}
}
}
return sd, nil
}
func decTrickle(data []byte) (trickleMsg, error) {
r := pbReader{buf: data}
var m trickleMsg
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return m, err
}
switch {
case field == trickleFieldCandidate && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return m, err
}
m.CandidateInit = string(b)
case field == trickleFieldTarget && wire == wireVarint:
v, err := r.varint()
if err != nil {
return m, err
}
m.Target = int(v)
case field == trickleFieldFinal && wire == wireVarint:
v, err := r.varint()
if err != nil {
return m, err
}
m.Final = v != 0
default:
if err := r.skipWire(wire); err != nil {
return m, err
}
}
}
return m, nil
}
func decICEServer(data []byte) (iceServer, error) {
r := pbReader{buf: data}
var s iceServer
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return s, err
}
switch {
case field == iceServerFieldURLs && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return s, err
}
s.URLs = append(s.URLs, string(b))
case field == iceServerFieldUsername && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return s, err
}
s.Username = string(b)
case field == iceServerFieldCredential && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return s, err
}
s.Credential = string(b)
default:
if err := r.skipWire(wire); err != nil {
return s, err
}
}
}
return s, nil
}
func decRoom(data []byte) (string, string, error) {
r := pbReader{buf: data}
var sid, name string
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return "", "", err
}
switch {
case field == roomFieldSID && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return "", "", err
}
sid = string(b)
case field == roomFieldName && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return "", "", err
}
name = string(b)
default:
if err := r.skipWire(wire); err != nil {
return "", "", err
}
}
}
return sid, name, nil
}
func decParticipant(data []byte) (string, string, error) {
info := DecodeParticipantInfo(data)
return info.SID, info.Identity, nil
}
func decJoinResponse(data []byte) (joinResponse, error) {
r := pbReader{buf: data}
var jr joinResponse
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return jr, err
}
switch {
case field == joinFieldRoom && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return jr, err
}
sid, name, _ := decRoom(b)
jr.RoomSID = sid
jr.RoomName = name
case field == joinFieldParticipant && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return jr, err
}
sid, identity, _ := decParticipant(b)
jr.ParticipantSID = sid
jr.ParticipantID = identity
case field == joinFieldServerVersion && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return jr, err
}
jr.ServerVersion = string(b)
case field == joinFieldICEServers && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return jr, err
}
s, _ := decICEServer(b)
jr.ICEServers = append(jr.ICEServers, s)
case field == joinFieldSubscriberPrimary && wire == wireVarint:
v, err := r.varint()
if err != nil {
return jr, err
}
jr.SubscriberPrimary = v != 0
case field == joinFieldServerRegion && wire == wireBytes:
b, err := r.bytes()
if err != nil {
return jr, err
}
jr.ServerRegion = string(b)
case field == joinFieldPingTimeout && wire == wireVarint:
v, err := r.varint()
if err != nil {
return jr, err
}
jr.PingTimeoutSec = int32(v)
case field == joinFieldPingInterval && wire == wireVarint:
v, err := r.varint()
if err != nil {
return jr, err
}
jr.PingIntervalSec = int32(v)
case field == joinFieldOtherParticipants && wire == wireBytes:
if _, err := r.bytes(); err != nil {
return jr, err
}
default:
if err := r.skipWire(wire); err != nil {
return jr, err
}
}
}
return jr, nil
}
func decSignalResponse(data []byte) (signalResponse, error) {
r := pbReader{buf: data}
var sr signalResponse
for !r.eof() {
field, wire, err := r.tag()
if err != nil {
return sr, err
}
if wire != wireBytes {
if err := r.skipWire(wire); err != nil {
return sr, err
}
continue
}
inner, err := r.bytes()
if err != nil {
return sr, err
}
sr.Kind = int(field)
switch field {
case signalRespJoin:
jr, err := decJoinResponse(inner)
if err != nil {
return sr, err
}
sr.Join = &jr
case signalRespAnswer, signalRespOffer:
sd, err := decSessionDescription(inner)
if err != nil {
return sr, err
}
sr.SDP = &sd
case signalRespTrickle:
tm, err := decTrickle(inner)
if err != nil {
return sr, err
}
sr.Trickle = &tm
case signalRespRefreshToken:
sr.Token = string(inner)
case signalRespLeave:
li := DecodeLeaveRequest(inner)
sr.Leave = &li
case signalRespUpdate:
sr.Participants = DecodeParticipantUpdate(inner)
}
return sr, nil
}
return sr, nil
}

View File

@@ -0,0 +1,132 @@
package livekit
import "fmt"
const (
wireVarint = 0
wireFixed64 = 1
wireBytes = 2
wireFixed32 = 5
)
type pbWriter struct{ buf []byte }
func (w *pbWriter) varint(v uint64) {
for v >= 0x80 {
w.buf = append(w.buf, byte(v)|0x80)
v >>= 7
}
w.buf = append(w.buf, byte(v))
}
func (w *pbWriter) tag(field, wire uint64) { w.varint(field<<3 | wire) }
func (w *pbWriter) string(field uint64, s string) {
w.tag(field, wireBytes)
w.varint(uint64(len(s)))
w.buf = append(w.buf, s...)
}
func (w *pbWriter) bytes(field uint64, b []byte) {
w.tag(field, wireBytes)
w.varint(uint64(len(b)))
w.buf = append(w.buf, b...)
}
func (w *pbWriter) message(field uint64, b []byte) { w.bytes(field, b) }
func (w *pbWriter) int32(field uint64, v int32) {
w.tag(field, wireVarint)
w.varint(uint64(uint32(v)))
}
func (w *pbWriter) int64(field uint64, v int64) {
w.tag(field, wireVarint)
w.varint(uint64(v))
}
func (w *pbWriter) uint32(field uint64, v uint32) {
w.tag(field, wireVarint)
w.varint(uint64(v))
}
func (w *pbWriter) bool(field uint64, v bool) {
w.tag(field, wireVarint)
if v {
w.varint(1)
} else {
w.varint(0)
}
}
type pbReader struct {
buf []byte
pos int
}
func (r *pbReader) eof() bool { return r.pos >= len(r.buf) }
func (r *pbReader) varint() (uint64, error) {
var v uint64
var shift uint
for {
if r.pos >= len(r.buf) {
return 0, fmt.Errorf("varint: unexpected eof")
}
b := r.buf[r.pos]
r.pos++
v |= uint64(b&0x7f) << shift
if b < 0x80 {
return v, nil
}
shift += 7
if shift >= 64 {
return 0, fmt.Errorf("varint: overflow")
}
}
}
func (r *pbReader) tag() (field, wire uint64, err error) {
t, err := r.varint()
if err != nil {
return 0, 0, err
}
return t >> 3, t & 7, nil
}
func (r *pbReader) bytes() ([]byte, error) {
n, err := r.varint()
if err != nil {
return nil, err
}
if r.pos+int(n) > len(r.buf) {
return nil, fmt.Errorf("bytes: short read")
}
out := r.buf[r.pos : r.pos+int(n)]
r.pos += int(n)
return out, nil
}
func (r *pbReader) skipWire(wire uint64) error {
switch wire {
case wireVarint:
_, err := r.varint()
return err
case wireFixed64:
if r.pos+8 > len(r.buf) {
return fmt.Errorf("skip: short fixed64")
}
r.pos += 8
return nil
case wireBytes:
_, err := r.bytes()
return err
case wireFixed32:
if r.pos+4 > len(r.buf) {
return fmt.Errorf("skip: short fixed32")
}
r.pos += 4
return nil
}
return fmt.Errorf("unknown wire type %d", wire)
}

View File

@@ -0,0 +1,341 @@
package telemost
import (
"encoding/json"
"fmt"
"io"
mathrand "math/rand"
"net/http"
"net/url"
"strings"
"time"
"github.com/google/uuid"
"github.com/pion/interceptor"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing-box/transport/call/common"
)
const (
APIBase = "https://cloud-api.yandex.ru/telemost_front/v2/telemost"
Origin = "https://telemost.yandex.ru"
)
var CapabilitiesOffer = map[string][]string{
"offerAnswerMode": {"SEPARATE"},
"initialSubscriberOffer": {"ON_HELLO"},
"slotsMode": {"FROM_CONTROLLER"},
"simulcastMode": {"DISABLED", "STATIC"},
"selfVadStatus": {"FROM_SERVER", "FROM_CLIENT"},
"dataChannelSharing": {"TO_RTP"},
"videoEncoderConfig": {"NO_CONFIG", "ONLY_INIT_CONFIG", "RUNTIME_CONFIG"},
"dataChannelVideoCodec": {"VP8", "UNIQUE_CODEC_FROM_TRACK_DESCRIPTION"},
"bandwidthLimitationReason": {"BANDWIDTH_REASON_DISABLED", "BANDWIDTH_REASON_ENABLED"},
"sdkDefaultDeviceManagement": {"SDK_DEFAULT_DEVICE_MANAGEMENT_DISABLED", "SDK_DEFAULT_DEVICE_MANAGEMENT_ENABLED"},
"joinOrderLayout": {"JOIN_ORDER_LAYOUT_DISABLED", "JOIN_ORDER_LAYOUT_ENABLED"},
"pinLayout": {"PIN_LAYOUT_DISABLED"},
"sendSelfViewVideoSlot": {"SEND_SELF_VIEW_VIDEO_SLOT_DISABLED", "SEND_SELF_VIEW_VIDEO_SLOT_ENABLED"},
"serverLayoutTransition": {"SERVER_LAYOUT_TRANSITION_DISABLED"},
"sdkPublisherOptimizeBitrate": {"SDK_PUBLISHER_OPTIMIZE_BITRATE_DISABLED", "SDK_PUBLISHER_OPTIMIZE_BITRATE_FULL", "SDK_PUBLISHER_OPTIMIZE_BITRATE_ONLY_SELF"},
"sdkNetworkLostDetection": {"SDK_NETWORK_LOST_DETECTION_DISABLED"},
"sdkNetworkPathMonitor": {"SDK_NETWORK_PATH_MONITOR_DISABLED"},
"publisherVp9": {"PUBLISH_VP9_DISABLED", "PUBLISH_VP9_ENABLED"},
"svcMode": {"SVC_MODE_DISABLED", "SVC_MODE_L3T3", "SVC_MODE_L3T3_KEY"},
"subscriberOfferAsyncAck": {"SUBSCRIBER_OFFER_ASYNC_ACK_DISABLED", "SUBSCRIBER_OFFER_ASYNC_ACK_ENABLED"},
"subscriberDtlsPassiveMode": {"SUBSCRIBER_DTLS_PASSIVE_MODE_DISABLED", "SUBSCRIBER_DTLS_PASSIVE_MODE_ENABLED"},
"androidBluetoothRoutingFix": {"ANDROID_BLUETOOTH_ROUTING_FIX_DISABLED"},
"fixedIceCandidatesPoolSize": {"FIXED_ICE_CANDIDATES_POOL_SIZE_DISABLED"},
"sdkAndroidTelecomIntegration": {"SDK_ANDROID_TELECOM_INTEGRATION_DISABLED"},
"setActiveCodecsMode": {"SET_ACTIVE_CODECS_MODE_DISABLED", "SET_ACTIVE_CODECS_MODE_VIDEO_ONLY"},
"publisherOpusDred": {"PUBLISHER_OPUS_DRED_DISABLED"},
"publisherOpusLowBitrate": {"PUBLISHER_OPUS_LOW_BITRATE_DISABLED"},
"sdkAndroidDestroySessionOnTaskRemoved": {"SDK_ANDROID_DESTROY_SESSION_ON_TASK_REMOVED_DISABLED"},
"svcModes": {"FALSE"},
"reportTelemetryModes": {"TRUE"},
"keepDefaultDevicesModes": {"FALSE"},
}
var StartupSlotSizes = [][][2]int{
{{0, 0}, {0, 0}, {0, 0}, {0, 0}, {0, 0}, {0, 0}, {0, 0}, {0, 0}, {0, 0}, {0, 0}, {0, 0}, {0, 0}},
{{464, 261}, {464, 261}, {464, 261}, {336, 189}, {272, 153}, {272, 153}, {272, 153}, {272, 153}, {224, 126}, {224, 126}, {224, 126}, {224, 126}},
{{464, 261}, {464, 261}, {464, 261}, {336, 189}, {272, 153}, {272, 153}, {272, 153}, {272, 153}, {224, 126}, {224, 126}, {224, 126}, {224, 126}},
{{672, 378}, {672, 378}, {464, 261}, {336, 189}, {320, 180}, {320, 180}, {320, 180}, {320, 180}, {272, 153}, {272, 153}, {224, 126}, {224, 126}},
}
type SlotBindEvent struct {
Slot int
ParticipantID string
Mid string
Reason string
}
type Client struct {
HTTP *http.Client
Cookie string
UserAgent string
AppVersion string
InstanceID string
}
func (c *Client) Do(method, path string, body interface{}) ([]byte, int, error) {
var bodyReader io.Reader
if body != nil {
data, _ := json.Marshal(body)
bodyReader = strings.NewReader(string(data))
}
req, err := http.NewRequest(method, APIBase+path, bodyReader)
if err != nil {
return nil, 0, err
}
ua := c.UserAgent
if ua == "" {
ua = common.UserAgent
}
instanceID := c.InstanceID
if instanceID == "" {
instanceID = uuid.New().String()
}
req.Header.Set("User-Agent", ua)
req.Header.Set("Origin", Origin)
req.Header.Set("Referer", Origin+"/")
req.Header.Set("Client-Instance-Id", instanceID)
if c.Cookie != "" {
req.Header.Set("Cookie", c.Cookie)
}
if c.AppVersion != "" {
req.Header.Set("X-Telemost-Client-Version", c.AppVersion)
}
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
client := c.HTTP
if client == nil {
client = http.DefaultClient
}
resp, err := client.Do(req)
if err != nil {
return nil, 0, err
}
defer resp.Body.Close()
data, err := io.ReadAll(resp.Body)
return data, resp.StatusCode, err
}
func (c *Client) TMRequest(method, path string) ([]byte, int, error) {
return c.Do(method, path, nil)
}
func (c *Client) RequestStates(joinURI, peerID string) error {
confURL := url.QueryEscape(joinURI)
body := map[string]interface{}{
"peers": []map[string]string{{"peer_id": peerID}},
"permissions": map[string]interface{}{},
"conference": map[string]interface{}{"version": -1},
}
r, status, err := c.Do("POST", "/conferences/"+confURL+"/request-states", body)
if err != nil {
return err
}
if status != 200 {
return fmt.Errorf("status %d: %s", status, string(r))
}
return nil
}
func NewAPI(settingEngine *webrtc.SettingEngine) (*webrtc.API, error) {
mediaEngine := &webrtc.MediaEngine{}
if err := mediaEngine.RegisterDefaultCodecs(); err != nil {
return nil, err
}
for _, uri := range []string{
"urn:ietf:params:rtp-hdrext:toffset",
"http://www.webrtc.org/experiments/rtp-hdrext/abs-send-time",
"urn:3gpp:video-orientation",
"http://www.webrtc.org/experiments/rtp-hdrext/playout-delay",
"http://www.webrtc.org/experiments/rtp-hdrext/video-content-type",
"http://www.webrtc.org/experiments/rtp-hdrext/video-timing",
"http://www.webrtc.org/experiments/rtp-hdrext/color-space",
} {
if err := mediaEngine.RegisterHeaderExtension(
webrtc.RTPHeaderExtensionCapability{URI: uri},
webrtc.RTPCodecTypeVideo,
); err != nil {
return nil, fmt.Errorf("register header extension %s: %w", uri, err)
}
}
registry := &interceptor.Registry{}
if err := webrtc.RegisterDefaultInterceptors(mediaEngine, registry); err != nil {
return nil, err
}
opts := []func(*webrtc.API){
webrtc.WithMediaEngine(mediaEngine),
webrtc.WithInterceptorRegistry(registry),
}
if settingEngine != nil {
opts = append(opts, webrtc.WithSettingEngine(*settingEngine))
}
return webrtc.NewAPI(opts...), nil
}
func NewPeerConnection(config webrtc.Configuration) (*webrtc.PeerConnection, error) {
api, err := NewAPI(nil)
if err != nil {
return nil, err
}
return api.NewPeerConnection(config)
}
func MungeSDPAddVideoContent(sdp string) string {
lines := strings.Split(sdp, "\r\n")
out := make([]string, 0, len(lines)+4)
inVideo := false
inserted := false
for _, line := range lines {
if strings.HasPrefix(line, "m=") {
if inVideo && !inserted {
out = append(out, "a=content:speaker,main")
inserted = true
}
inVideo = strings.HasPrefix(line, "m=video")
inserted = false
}
out = append(out, line)
if inVideo && !inserted && strings.HasPrefix(line, "a=mid:") {
out = append(out, "a=content:speaker,main")
inserted = true
}
}
return strings.Join(out, "\r\n")
}
func SlotsConfigBindings(v interface{}) []SlotBindEvent {
m, ok := v.(map[string]interface{})
if !ok {
return nil
}
slots, _ := m["slots"].([]interface{})
var out []SlotBindEvent
for idx, s := range slots {
sm, _ := s.(map[string]interface{})
if pv, _ := sm["participantVideoByMid"].(map[string]interface{}); pv != nil {
pid, _ := pv["participantId"].(string)
mid, _ := pv["mid"].(string)
reason, _ := pv["limitationReason"].(string)
out = append(out, SlotBindEvent{Slot: idx, ParticipantID: pid, Mid: mid, Reason: reason})
continue
}
if p, _ := sm["participant"].(map[string]interface{}); p != nil {
pid, _ := p["participantId"].(string)
out = append(out, SlotBindEvent{Slot: idx, ParticipantID: pid})
}
}
return out
}
func BriefJSON(v interface{}) string {
const max = 240
b, err := json.Marshal(v)
if err != nil {
return fmt.Sprintf("<json err: %v>", err)
}
if len(b) > max {
return string(b[:max]) + "...(+" + fmt.Sprintf("%d", len(b)-max) + "B)"
}
return string(b)
}
func SetSlotsMessage(key int) map[string]interface{} {
rnd := mathrand.New(mathrand.NewSource(time.Now().UnixNano()))
return slotsMessageWithSizes(key, StartupSlotSizes[len(StartupSlotSizes)-1], rnd)
}
func StartupSetSlotsMessage(i, key int) map[string]interface{} {
rnd := mathrand.New(mathrand.NewSource(time.Now().UnixNano() + int64(i)))
return slotsMessageWithSizes(key, StartupSlotSizes[i], rnd)
}
func SetSlotsOffsetMessage(offset int) map[string]interface{} {
return map[string]interface{}{
"uid": uuid.New().String(),
"setSlotsOffset": map[string]interface{}{"offset": offset},
}
}
func SdkCodecsInfoMessage() map[string]interface{} {
return map[string]interface{}{
"uid": uuid.New().String(),
"sdkCodecsInfo": map[string]interface{}{
"vp8": map[string]interface{}{
"supported": "CODEC_FEATURE_SUPPORTED",
"hwDecode": "CODEC_FEATURE_NOT_SUPPORTED",
"hwEncode": "CODEC_FEATURE_NOT_SUPPORTED",
"isoString": "vp8",
},
},
}
}
func UpdatePublisherTrackDescriptionMessage(pc *webrtc.PeerConnection, audioLabel, videoLabel string) map[string]interface{} {
descs := []map[string]interface{}{}
for _, tr := range pc.GetTransceivers() {
sender := tr.Sender()
if sender == nil || sender.Track() == nil {
continue
}
kind := strings.ToUpper(sender.Track().Kind().String())
mid := tr.Mid()
label := videoLabel
groupId := 2
if kind == "AUDIO" {
label = audioLabel
groupId = 1
}
descs = append(descs, map[string]interface{}{
"mid": mid,
"transceiverMid": mid,
"kind": kind,
"priority": 0,
"label": label,
"codecs": map[string]interface{}{},
"groupId": groupId,
"description": "",
})
}
return map[string]interface{}{
"uid": uuid.New().String(),
"updatePublisherTrackDescription": map[string]interface{}{
"publisherTrackDescriptions": descs,
},
}
}
func jitterSize(width int, rnd *mathrand.Rand) (int, int) {
if width == 0 {
return 0, 0
}
w := width + rnd.Intn(11) - 5
return w, w * 9 / 16
}
func slotsMessageWithSizes(key int, template [][2]int, rnd *mathrand.Rand) map[string]interface{} {
slots := make([]map[string]interface{}, len(template))
for i, wh := range template {
w, h := wh[0], wh[1]
if rnd != nil {
w, h = jitterSize(wh[0], rnd)
}
slots[i] = map[string]interface{}{"width": w, "height": h}
}
return map[string]interface{}{
"uid": uuid.New().String(),
"setSlots": map[string]interface{}{
"slots": slots,
"audioSlotsCount": 0,
"key": key,
"shutdownAllVideo": nil,
"withSelfView": true,
"selfViewVisibility": "ON_LOADING_THEN_SHOW",
"gridConfig": map[string]interface{}{},
},
}
}

View File

@@ -0,0 +1,72 @@
package telemost
import (
"encoding/json"
"fmt"
"regexp"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
type TMConfig struct {
AppVersion string
SDKVersion string
}
func FetchConfig(dialer N.Dialer, logger logger.ContextLogger) (TMConfig, error) {
var cfg TMConfig
page, err := common.HttpGet(dialer, "https://telemost.yandex.ru/")
if err != nil {
return cfg, fmt.Errorf("failed to fetch telemost.yandex.ru: %w", err)
}
stateRe := regexp.MustCompile(`<script[^>]*id="preloaded-state"[^>]*>([\s\S]*?)</script>`)
stateMatch := stateRe.FindSubmatch(page)
if stateMatch == nil {
return cfg, fmt.Errorf("preloaded-state not found in page")
}
var state struct {
Config struct {
AppVersion string `json:"appVersion"`
} `json:"config"`
AppVersion string `json:"appVersion"`
}
if err := json.Unmarshal(stateMatch[1], &state); err != nil {
return cfg, fmt.Errorf("failed to parse preloaded-state: %w", err)
}
cfg.AppVersion = state.Config.AppVersion
if cfg.AppVersion == "" {
cfg.AppVersion = state.AppVersion
}
if cfg.AppVersion == "" {
return cfg, fmt.Errorf("appVersion not found in preloaded-state")
}
logger.Debug(fmt.Sprintf("[config] appVersion=%s", cfg.AppVersion))
bundleRe := regexp.MustCompile(`https://telemost\.yastatic\.net/s3/telemost/_/main\.\w+\.[a-f0-9]+\.js`)
bundleURL := bundleRe.FindString(string(page))
if bundleURL == "" {
return cfg, fmt.Errorf("main bundle URL not found in page")
}
logger.Debug(fmt.Sprintf("[config] Found bundle: %s", bundleURL))
bundle, err := common.HttpGet(dialer, bundleURL)
if err != nil {
return cfg, fmt.Errorf("failed to fetch bundle: %w", err)
}
sdkVerPatterns := []*regexp.Regexp{
regexp.MustCompile(`goloom_sdk_version:"(\d+\.\d+\.\d+)"`),
regexp.MustCompile(`"@yandex-video-platform/goloom-sdk":"(\d+\.\d+\.\d+)"`),
regexp.MustCompile(`goloom-sdk\.(\d+\.\d+\.\d+)\.js`),
}
for _, re := range sdkVerPatterns {
if m := re.FindSubmatch(bundle); m != nil {
cfg.SDKVersion = string(m[1])
break
}
}
if cfg.SDKVersion == "" {
return cfg, fmt.Errorf("goloom SDK version not found in bundle")
}
logger.Debug(fmt.Sprintf("[config] app=%s sdk=%s", cfg.AppVersion, cfg.SDKVersion))
return cfg, nil
}

View File

@@ -0,0 +1,97 @@
package telemost
import (
"context"
"fmt"
"time"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing-box/transport/call/tunnel"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
func ConnectCreator(ctx context.Context, cookieStr, joinLink string, readBuf int, dialer N.Dialer, logger logger.ContextLogger) (*tunnel.RelayBridge, string, error) {
cfg, err := FetchConfig(dialer, logger)
if err != nil {
return nil, "", err
}
var connInfo *ConnInfo
if joinLink != "" {
connInfo, err = joinExistingConference(dialer, cookieStr, joinLink, cfg, logger)
} else {
connInfo, err = CreateAndJoinCall(dialer, cookieStr, cfg, logger)
}
if err != nil {
return nil, "", err
}
if readBuf <= 0 {
readBuf = 32768
}
bridge := &Bridge{
connInfo: connInfo,
config: cfg,
cookieStr: cookieStr,
peers: make(map[string]string),
readBuf: readBuf,
dialer: dialer,
logger: logger,
}
go bridge.Run()
deadline := time.Now().Add(60 * time.Second)
for bridge.activeBridge == nil {
if time.Now().After(deadline) {
return nil, "", fmt.Errorf("telemost: creator tunnel timed out")
}
select {
case <-ctx.Done():
return nil, "", ctx.Err()
case <-time.After(200 * time.Millisecond):
}
}
return bridge.activeBridge, connInfo.ConferenceURI, nil
}
func ConnectJoiner(ctx context.Context, joinLink, displayName string, readBuf int, dialer N.Dialer, dnsRouter adapter.DNSRouter, logger logger.ContextLogger) (tunnel.DataTunnel, error) {
if displayName == "" {
displayName = "Joiner"
}
joiner := NewTelemostJoiner(
logger,
dialer,
dnsRouter,
nil,
common.AddTunnelTracks,
common.ReadTrack,
)
tunCh := make(chan tunnel.DataTunnel, 1)
joiner.OnConnected = func(tun tunnel.DataTunnel) {
select {
case tunCh <- tun:
default:
}
}
params := fmt.Sprintf(`{"joinLink":%q,"displayName":%q}`, joinLink, displayName)
go joiner.RunWithParams(params)
select {
case tun := <-tunCh:
return tun, nil
case <-ctx.Done():
joiner.Close()
return nil, ctx.Err()
}
}
func CreateConferenceForTest(dialer N.Dialer, cookieStr string) (string, error) {
nop := logger.NOP()
cfg, err := FetchConfig(dialer, nop)
if err != nil {
return "", err
}
connInfo, err := CreateAndJoinCall(dialer, cookieStr, cfg, nop)
if err != nil {
return "", err
}
return connInfo.ConferenceURI, nil
}

View File

@@ -0,0 +1,848 @@
package telemost
import (
"context"
"encoding/json"
"fmt"
"net"
"net/http"
"net/url"
"strings"
"sync"
"time"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"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"
)
const (
tmAPIBase = APIBase
tmOrigin = Origin
tmPingPeriod = 5 * time.Second
)
var clientInstanceID = uuid.New().String()
type ConnInfo struct {
ConferenceURI string
RoomID string
PeerID string
Credentials string
MediaServerURL string
ServiceName string
ICEServers []webrtc.ICEServer
StateCheckIntervalS int
}
type Bridge struct {
mu sync.Mutex
ws *websocket.Conn
relay *SFURelay
connInfo *ConnInfo
config TMConfig
cookieStr string
pubSeq int
subSeq int
peers map[string]string
readBuf int
activeBridge *tunnel.RelayBridge
selfName string
dialer N.Dialer
logger logger.ContextLogger
setSlotsKey int
initBundleSent bool
pendingKicks map[string]chan struct{}
boundPeers map[string]bool
unboundPeers map[string]bool
}
func tmRequest(dialer N.Dialer, method, path string, body interface{}, cookieStr string, cfg TMConfig) ([]byte, int, error) {
c := Client{HTTP: common.HttpClient(dialer), Cookie: cookieStr, AppVersion: cfg.AppVersion, InstanceID: clientInstanceID}
return c.Do(method, path, body)
}
func parseICEServersJSON(raw json.RawMessage) []webrtc.ICEServer {
var rawIce []struct {
URLs []string `json:"urls"`
Username string `json:"username"`
Credential string `json:"credential"`
}
json.Unmarshal(raw, &rawIce)
var out []webrtc.ICEServer
for _, s := range rawIce {
ice := webrtc.ICEServer{URLs: s.URLs}
if s.Username != "" {
ice.Username = s.Username
ice.Credential = s.Credential
}
out = append(out, ice)
}
return out
}
func getConnection(dialer N.Dialer, cookieStr, confURL string, cfg TMConfig) (*ConnInfo, error) {
r, status, err := tmRequest(dialer, "GET",
"/conferences/"+confURL+"/connection?next_gen_media_platform_allowed=true&display_name=Headless&waiting_room_supported=true",
nil, cookieStr, cfg)
if err != nil {
return nil, fmt.Errorf("get connection: %w", err)
}
if status != 200 {
return nil, fmt.Errorf("get connection: status %d: %s", status, string(r))
}
var conn struct {
PeerID string `json:"peer_id"`
RoomID string `json:"room_id"`
Credentials string `json:"credentials"`
ClientConfig struct {
MediaServerURL string `json:"media_server_url"`
ServiceName string `json:"service_name"`
ICEServers json.RawMessage `json:"ice_servers"`
StateCheckIntervalSecs int `json:"state_check_interval_seconds"`
} `json:"client_configuration"`
}
json.Unmarshal(r, &conn)
if conn.ClientConfig.MediaServerURL == "" {
return nil, fmt.Errorf("empty media_server_url: %s", string(r))
}
return &ConnInfo{
RoomID: conn.RoomID,
PeerID: conn.PeerID,
Credentials: conn.Credentials,
MediaServerURL: conn.ClientConfig.MediaServerURL,
ServiceName: conn.ClientConfig.ServiceName,
ICEServers: parseICEServersJSON(conn.ClientConfig.ICEServers),
StateCheckIntervalS: conn.ClientConfig.StateCheckIntervalSecs,
}, nil
}
func joinExistingConference(dialer N.Dialer, cookieStr, conferenceURI string, cfg TMConfig, logger logger.ContextLogger) (*ConnInfo, error) {
conferenceURI = strings.TrimSpace(conferenceURI)
if conferenceURI == "" {
return nil, fmt.Errorf("empty -tm-link")
}
logger.Info(fmt.Sprintf("[auth] Joining existing conference: %s", conferenceURI))
info, err := getConnection(dialer, cookieStr, url.QueryEscape(conferenceURI), cfg)
if err != nil {
return nil, err
}
info.ConferenceURI = conferenceURI
logger.Debug(fmt.Sprintf("[auth] peer_id=%s room_id=%s", info.PeerID, info.RoomID))
logger.Debug(fmt.Sprintf("[auth] media_server=%s", info.MediaServerURL))
return info, nil
}
func CreateAndJoinCall(dialer N.Dialer, cookieStr string, cfg TMConfig, logger logger.ContextLogger) (*ConnInfo, error) {
logger.Info("[auth] Creating conference...")
r, status, err := tmRequest(dialer, "POST", "/conferences?next_gen_media_platform_allowed=true",
struct{}{}, cookieStr, cfg)
if err != nil {
return nil, fmt.Errorf("create conference: %w", err)
}
if status != 200 && status != 201 {
return nil, fmt.Errorf("create conference: status %d: %s", status, string(r))
}
var conf struct {
URI string `json:"uri"`
}
json.Unmarshal(r, &conf)
if conf.URI == "" {
return nil, fmt.Errorf("empty conference URI: %s", string(r))
}
logger.Info(fmt.Sprintf("[auth] Conference: %s", conf.URI))
logger.Debug("[auth] Getting connection...")
info, err := getConnection(dialer, cookieStr, url.QueryEscape(conf.URI), cfg)
if err != nil {
return nil, err
}
info.ConferenceURI = conf.URI
logger.Debug(fmt.Sprintf("[auth] peer_id=%s room_id=%s", info.PeerID, info.RoomID))
logger.Debug(fmt.Sprintf("[auth] media_server=%s", info.MediaServerURL))
return info, nil
}
func (b *Bridge) wsSend(msg interface{}) {
b.mu.Lock()
defer b.mu.Unlock()
if b.ws == nil {
return
}
data, _ := json.Marshal(msg)
b.ws.WriteMessage(websocket.TextMessage, data)
}
func (b *Bridge) ack(uid string) {
b.wsSend(map[string]interface{}{
"uid": uid,
"ack": map[string]interface{}{
"status": map[string]interface{}{"code": "OK", "description": ""},
},
})
}
func (b *Bridge) sendHello() {
b.mu.Lock()
b.selfName = "Headless"
b.mu.Unlock()
b.wsSend(map[string]interface{}{
"uid": uuid.New().String(),
"hello": map[string]interface{}{
"participantMeta": map[string]interface{}{"name": "Headless", "role": "SPEAKER", "description": "", "sendAudio": false, "sendVideo": true},
"participantAttributes": map[string]interface{}{"name": "Headless", "role": "SPEAKER", "description": ""},
"sendAudio": false, "sendVideo": true, "sendSharing": false,
"participantId": b.connInfo.PeerID, "roomId": b.connInfo.RoomID,
"serviceName": b.connInfo.ServiceName, "credentials": b.connInfo.Credentials,
"capabilitiesOffer": CapabilitiesOffer,
"sdkInfo": map[string]interface{}{"implementation": "browser", "version": b.config.SDKVersion, "userAgent": common.UserAgent, "hwConcurrency": 8},
"sdkInitializationId": uuid.New().String(),
"disablePublisher": false, "disableSubscriber": false, "disableSubscriberAudio": false,
},
})
b.logger.Debug("[tm-ws] -> hello")
}
func (b *Bridge) sendPubOffer() {
offer, err := b.relay.CreatePubOffer()
if err != nil {
b.logger.Warn(fmt.Sprintf("[tm-ws] pub offer failed: %v", err))
return
}
audioMid, videoMid := parseMids(offer.SDP)
b.logger.Debug(fmt.Sprintf("[tm-ws] -> publisherSdpOffer pcSeq=%d", b.pubSeq))
var tracks []map[string]interface{}
if audioMid != "" {
tracks = append(tracks, map[string]interface{}{"mid": audioMid, "transceiverMid": audioMid, "kind": "AUDIO", "priority": 0, "label": "", "codecs": map[string]interface{}{}, "groupId": 1, "description": ""})
}
if videoMid != "" {
tracks = append(tracks, map[string]interface{}{"mid": videoMid, "transceiverMid": videoMid, "kind": "VIDEO", "priority": 0, "label": "", "codecs": map[string]interface{}{}, "groupId": 2, "description": ""})
}
b.wsSend(map[string]interface{}{
"uid": uuid.New().String(),
"publisherSdpOffer": map[string]interface{}{"pcSeq": b.pubSeq, "sdp": offer.SDP, "tracks": tracks},
})
}
func (b *Bridge) sendICE(cand *webrtc.ICECandidate, target string, pcSeq int) {
c := cand.ToJSON()
mid := ""
if c.SDPMid != nil {
mid = *c.SDPMid
}
var idx uint16
if c.SDPMLineIndex != nil {
idx = *c.SDPMLineIndex
}
b.wsSend(map[string]interface{}{
"uid": uuid.New().String(),
"webrtcIceCandidate": map[string]interface{}{
"candidate": c.Candidate, "sdpMid": mid,
"usernameFragment": extractUfrag(c.Candidate),
"sdpMlineIndex": idx, "target": target, "pcSeq": pcSeq,
},
})
}
func (b *Bridge) requestVideoSlots() {
b.setSlotsKey++
b.logger.Debug(fmt.Sprintf("[tm-ws] -> setSlots key=%d", b.setSlotsKey))
b.wsSend(SetSlotsMessage(b.setSlotsKey))
}
func (b *Bridge) forceReconnect(reason string) {
oldPeerID := b.connInfo.PeerID
b.logger.Info(fmt.Sprintf("[tm-ws] forcing reconnect: %s", reason))
if oldPeerID != "" {
b.logger.Debug(fmt.Sprintf("[tm-ws] kicking self pid=%s to leave call cleanly", oldPeerID))
if err := b.kickPeer(oldPeerID); err != nil {
b.logger.Warn(fmt.Sprintf("[tm-ws] self-kick failed: %v", err))
}
}
clientInstanceID = uuid.New().String()
b.logger.Debug(fmt.Sprintf("[tm-ws] new instance-id=%s", clientInstanceID))
b.mu.Lock()
ws := b.ws
b.mu.Unlock()
if ws != nil {
ws.Close()
}
}
func (b *Bridge) sendInitBundle() {
if b.initBundleSent {
return
}
b.initBundleSent = true
b.logger.Debug("[tm-ws] -> sdkCodecsInfo + updatePublisherTrackDescription")
b.wsSend(SdkCodecsInfoMessage())
b.wsSend(UpdatePublisherTrackDescriptionMessage(b.relay.pubPC, "Microphone", "MacBook Pro Camera (0000:0001)"))
b.sendStartupSlotsRamp()
}
func (b *Bridge) sendStartupSlotsRamp() {
for i := 0; i < 4; i++ {
b.setSlotsKey++
b.logger.Debug(fmt.Sprintf("[tm-ws] -> setSlots key=%d (startup %d/4)", b.setSlotsKey, i+1))
b.wsSend(StartupSetSlotsMessage(i, b.setSlotsKey))
}
}
func (b *Bridge) handleMessage(raw []byte) {
var msg map[string]interface{}
if err := json.Unmarshal(raw, &msg); err != nil {
return
}
uid, _ := msg["uid"].(string)
if sh, ok := msg["serverHello"]; ok {
b.logger.Debug("[tm-ws] <- serverHello")
if shMap, ok := sh.(map[string]interface{}); ok {
b.parseICEServers(shMap)
}
b.ack(uid)
b.logger.Debug("[tm-ws] -> setSlotsOffset")
b.wsSend(SetSlotsOffsetMessage(0))
b.initRelay()
return
}
if pa, ok := msg["publisherSdpAnswer"]; ok {
paMap, _ := pa.(map[string]interface{})
sdp, _ := paMap["sdp"].(string)
b.logger.Debug(fmt.Sprintf("[tm-ws] <- publisherSdpAnswer %d bytes", len(sdp)))
if err := b.relay.SetPubAnswer(sdp); err != nil {
b.logger.Warn(fmt.Sprintf("[tm-ws] error: %v", err))
return
}
b.sendInitBundle()
return
}
if so, ok := msg["subscriberSdpOffer"]; ok {
soMap, _ := so.(map[string]interface{})
sdp, _ := soMap["sdp"].(string)
pcSeq, _ := soMap["pcSeq"].(float64)
b.subSeq = int(pcSeq)
b.logger.Debug(fmt.Sprintf("[tm-ws] <- subscriberSdpOffer pcSeq=%d", b.subSeq))
b.ack(uid)
answer, err := b.relay.SetSubOffer(sdp)
if err != nil {
b.logger.Warn(fmt.Sprintf("[tm-ws] error: %v", err))
return
}
b.logger.Debug(fmt.Sprintf("[tm-ws] -> subscriberSdpAnswer pcSeq=%d", b.subSeq))
b.wsSend(map[string]interface{}{
"uid": uuid.New().String(),
"subscriberSdpAnswer": map[string]interface{}{"sdp": answer.SDP, "pcSeq": b.subSeq},
})
b.sendPubOffer()
return
}
if ic, ok := msg["webrtcIceCandidate"]; ok {
icMap, _ := ic.(map[string]interface{})
candidate, _ := icMap["candidate"].(string)
sdpMid, _ := icMap["sdpMid"].(string)
target, _ := icMap["target"].(string)
sdpIdx, _ := icMap["sdpMlineIndex"].(float64)
idx := uint16(sdpIdx)
cand := webrtc.ICECandidateInit{Candidate: candidate, SDPMid: &sdpMid, SDPMLineIndex: &idx}
if target == "PUBLISHER" {
b.relay.AddPubICECandidate(cand)
} else {
b.relay.AddSubICECandidate(cand)
}
b.ack(uid)
return
}
if ackData, ok := msg["ack"]; ok {
if ackMap, ok := ackData.(map[string]interface{}); ok {
if status, ok := ackMap["status"].(map[string]interface{}); ok {
if code, _ := status["code"].(string); code != "OK" {
desc, _ := status["description"].(string)
b.logger.Warn(fmt.Sprintf("[tm-ws] <- ack error: %s %s", code, desc))
}
}
}
return
}
if ud, ok := msg["updateDescription"]; ok {
b.logger.Debug(fmt.Sprintf("[tm-ws] <- updateDescription %s", BriefJSON(ud)))
udMap, _ := ud.(map[string]interface{})
descs, _ := udMap["description"].([]interface{})
b.applyDescriptionSnapshot(descs)
b.ack(uid)
return
}
if ud, ok := msg["upsertDescription"]; ok {
udMap, _ := ud.(map[string]interface{})
descs, _ := udMap["description"].([]interface{})
for _, d := range descs {
dm, _ := d.(map[string]interface{})
b.applyDescriptionEntry(dm)
}
b.kickStaleSelves()
b.ack(uid)
return
}
if rd, ok := msg["removeDescription"]; ok {
rdMap, _ := rd.(map[string]interface{})
ids, _ := rdMap["descriptionId"].([]interface{})
for _, id := range ids {
pid, _ := id.(string)
b.mu.Lock()
name := b.peers[pid]
delete(b.peers, pid)
remaining := len(b.peers)
ch, hadPendingKick := b.pendingKicks[pid]
if hadPendingKick {
delete(b.pendingKicks, pid)
}
b.mu.Unlock()
b.logger.Info(fmt.Sprintf("[tm-ws] Participant left: %s (%s) total=%d", name, pid, remaining))
if hadPendingKick {
close(ch)
}
if remaining == 0 {
go b.pollAndAdmit()
}
}
b.ack(uid)
return
}
if n, ok := msg["notification"]; ok {
b.logger.Debug(fmt.Sprintf("[tm-ws] <- notification %s", BriefJSON(n)))
b.ack(uid)
go b.pollAndAdmit()
return
}
if pc, ok := msg["participantsChanged"]; ok {
b.logger.Debug(fmt.Sprintf("[tm-ws] <- participantsChanged %s", BriefJSON(pc)))
b.ack(uid)
go b.pollAndAdmit()
return
}
if sc, ok := msg["slotsConfig"]; ok {
b.logger.Debug(fmt.Sprintf("[tm-ws] <- slotsConfig %s", BriefJSON(sc)))
needRebind := false
presentPids := make(map[string]bool)
for _, ev := range SlotsConfigBindings(sc) {
fullPid := ev.ParticipantID
if fullPid != "" {
presentPids[fullPid] = true
}
pid := fullPid
if len(pid) > 8 {
pid = pid[:8]
}
if ev.Reason == "NO_LIMITATION" && ev.Mid != "" {
b.logger.Debug(fmt.Sprintf("[bind] BOUND slot=%d pid=%s mid=%s", ev.Slot, pid, ev.Mid))
b.mu.Lock()
if b.boundPeers == nil {
b.boundPeers = make(map[string]bool)
}
b.boundPeers[fullPid] = true
delete(b.unboundPeers, fullPid)
b.mu.Unlock()
} else if fullPid != "" {
b.mu.Lock()
wasBound := b.boundPeers[fullPid]
if wasBound {
if b.unboundPeers == nil {
b.unboundPeers = make(map[string]bool)
}
b.unboundPeers[fullPid] = true
delete(b.boundPeers, fullPid)
}
b.mu.Unlock()
if wasBound {
b.logger.Debug(fmt.Sprintf("[bind] KILL slot=%d pid=%s reason=%s - rebinding", ev.Slot, pid, ev.Reason))
needRebind = true
} else {
b.logger.Debug(fmt.Sprintf("[bind] UNBOUND slot=%d pid=%s reason=%s mid=%q", ev.Slot, pid, ev.Reason, ev.Mid))
}
}
}
b.mu.Lock()
for boundPid := range b.boundPeers {
if !presentPids[boundPid] {
short := boundPid
if len(short) > 8 {
short = short[:8]
}
b.logger.Debug(fmt.Sprintf("[bind] VANISHED pid=%s - rebinding", short))
delete(b.boundPeers, boundPid)
needRebind = true
}
}
b.mu.Unlock()
if needRebind {
go b.forceReconnect("slot binding killed")
}
b.ack(uid)
return
}
for k, v := range msg {
if k == "uid" || k == "ack" {
continue
}
b.logger.Debug(fmt.Sprintf("[tm-ws] <- %s (unhandled) %s", k, BriefJSON(v)))
break
}
if uid != "" {
b.ack(uid)
}
}
func (b *Bridge) parseICEServers(sh map[string]interface{}) {
rtcCfg, ok := sh["rtcConfiguration"].(map[string]interface{})
if !ok {
return
}
servers, ok := rtcCfg["iceServers"].([]interface{})
if !ok {
return
}
var iceServers []webrtc.ICEServer
for _, s := range servers {
sm, _ := s.(map[string]interface{})
var urls []string
if u, ok := sm["urls"].([]interface{}); ok {
for _, v := range u {
if vs, ok := v.(string); ok {
urls = append(urls, vs)
}
}
}
ice := webrtc.ICEServer{URLs: urls}
if u, ok := sm["username"].(string); ok && u != "" {
ice.Username = u
ice.Credential, _ = sm["credential"].(string)
}
iceServers = append(iceServers, ice)
}
b.connInfo.ICEServers = iceServers
b.logger.Debug(fmt.Sprintf("[tm-ws] %d ICE servers", len(iceServers)))
}
func (b *Bridge) requestStates() error {
c := Client{HTTP: common.HttpClient(b.dialer), Cookie: b.cookieStr, AppVersion: b.config.AppVersion, InstanceID: clientInstanceID}
return c.RequestStates(b.connInfo.ConferenceURI, b.connInfo.PeerID)
}
func (b *Bridge) applyDescriptionEntry(dm map[string]interface{}) {
pid, _ := dm["id"].(string)
if pid == "" {
return
}
name := ""
if meta, ok := dm["meta"].(map[string]interface{}); ok {
name, _ = meta["name"].(string)
}
if pid == b.connInfo.PeerID {
b.mu.Lock()
if name != "" {
b.selfName = name
}
b.mu.Unlock()
return
}
_, disconnected := dm["disconnectedAt"]
b.mu.Lock()
_, wasKnown := b.peers[pid]
if disconnected {
delete(b.peers, pid)
} else {
b.peers[pid] = name
}
total := len(b.peers)
b.mu.Unlock()
switch {
case disconnected && wasKnown:
b.logger.Info(fmt.Sprintf("[tm-ws] Participant left: %s (%s) total=%d", name, pid, total))
case disconnected:
b.logger.Debug(fmt.Sprintf("[tm-ws] Ghost participant: %s (%s) - kicking", name, pid))
go b.kickPeer(pid)
case !wasKnown:
b.logger.Info(fmt.Sprintf("[tm-ws] Participant joined: %s (%s) total=%d", name, pid, total))
}
}
func (b *Bridge) applyDescriptionSnapshot(descs []interface{}) {
b.mu.Lock()
b.peers = make(map[string]string)
b.mu.Unlock()
for _, d := range descs {
dm, _ := d.(map[string]interface{})
b.applyDescriptionEntry(dm)
}
b.kickStaleSelves()
}
func (b *Bridge) kickStaleSelves() {
b.mu.Lock()
selfName := b.selfName
stale := make([]string, 0)
if selfName != "" {
for pid, name := range b.peers {
if name == selfName {
stale = append(stale, pid)
}
}
}
b.mu.Unlock()
for _, pid := range stale {
b.logger.Debug(fmt.Sprintf("[tm-ws] Kicking stale self %s (name=%q)", pid, selfName))
b.kickPeer(pid)
b.mu.Lock()
delete(b.peers, pid)
b.mu.Unlock()
}
}
func (b *Bridge) kickPeer(peerID string) error {
confURL := url.QueryEscape(b.connInfo.ConferenceURI)
b.logger.Debug(fmt.Sprintf("[tm-ws] Kicking %s", peerID))
body, status, err := tmRequest(b.dialer, "POST", "/conferences/"+confURL+"/commands/kick?peer_id="+url.QueryEscape(peerID)+"&with_ban=false",
nil, b.cookieStr, b.config)
if err != nil {
return err
}
if status >= 400 {
return fmt.Errorf("kick %s status %d: %s", peerID, status, string(body))
}
return nil
}
func (b *Bridge) pollAndAdmit() {
confURL := url.QueryEscape(b.connInfo.ConferenceURI)
r, status, err := tmRequest(b.dialer, "GET", "/conferences/"+confURL+"/waiting-rooms/peers", nil, b.cookieStr, b.config)
if err != nil || status != 200 {
return
}
var resp struct {
Peers []struct {
PeerID string `json:"peer_id"`
State struct {
DisplayName string `json:"display_name"`
} `json:"state"`
} `json:"peers"`
}
json.Unmarshal(r, &resp)
if len(resp.Peers) == 0 {
return
}
b.mu.Lock()
if b.pendingKicks == nil {
b.pendingKicks = make(map[string]chan struct{})
}
toKick := make(map[string]string, len(b.peers))
waits := make(map[string]<-chan struct{}, len(b.peers))
for pid, name := range b.peers {
toKick[pid] = name
ch := make(chan struct{})
b.pendingKicks[pid] = ch
waits[pid] = ch
}
b.mu.Unlock()
for pid, name := range toKick {
b.logger.Debug(fmt.Sprintf("[tm-ws] Kicking %s (%s) for one-to-one", name, pid))
if err := b.kickPeer(pid); err != nil {
b.logger.Warn(fmt.Sprintf("[tm-ws] kick failed: %v", err))
b.mu.Lock()
delete(b.pendingKicks, pid)
b.mu.Unlock()
return
}
}
for pid, ch := range waits {
<-ch
b.logger.Debug(fmt.Sprintf("[tm-ws] kick confirmed for %s", pid))
}
p := resp.Peers[0]
b.logger.Debug(fmt.Sprintf("[tm-ws] Admitting %s (%s)", p.State.DisplayName, p.PeerID))
tmRequest(b.dialer, "PUT", "/conferences/"+confURL+"/commands/admit?peer_id="+url.QueryEscape(p.PeerID),
nil, b.cookieStr, b.config)
}
func (b *Bridge) initRelay() {
if b.relay != nil {
b.relay.Close()
}
b.pubSeq = 1
b.subSeq = 0
b.initBundleSent = false
relay := NewSFURelay(b.logger)
relay.readBufSize = b.readBuf
obf, err := tunnel.NewTunnelObfuscator(tunnel.DeriveSecretFromJoinLink(b.connInfo.ConferenceURI))
if err != nil {
b.logger.Fatal(fmt.Sprintf("[relay] obfuscator init failed: %v", err))
}
relay.SetObfuscator(obf)
b.logger.Debug(fmt.Sprintf("[relay] obfuscator localEpoch=0x%08x", obf.LocalEpoch()))
relay.OnPubReady = func() {
b.logger.Debug("[relay] pub PC connected")
}
relay.OnConnected = func(tun *tunnel.VP8DataTunnel) {
if b.activeBridge != nil {
b.activeBridge.Reset()
}
b.activeBridge = tunnel.NewRelayBridge(tun, "creator", common.VP8BufSize, b.dialer, b.logger)
b.logger.Debug("[relay] tunnel connected")
}
relay.OnPeerRestart = func() {
if b.activeBridge != nil {
b.logger.Info("[relay] new peer detected, resetting relay bridge")
b.activeBridge.Reset()
}
}
relay.OnPubICE = func(cand *webrtc.ICECandidate) {
if cand == nil {
return
}
b.sendICE(cand, "PUBLISHER", b.pubSeq)
}
relay.OnSubICE = func(cand *webrtc.ICECandidate) {
if cand == nil {
return
}
b.sendICE(cand, "SUBSCRIBER", b.subSeq)
}
if err := relay.Init(b.connInfo.ICEServers); err != nil {
b.logger.Fatal(fmt.Sprintf("[relay] init failed: %v", err))
}
b.relay = relay
}
func (b *Bridge) Run() {
wsHeader := http.Header{}
wsHeader.Set("User-Agent", common.UserAgent)
wsHeader.Set("Origin", tmOrigin)
wsDialer := websocket.Dialer{
NetDialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
return b.dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
},
}
for {
b.logger.Debug("[tm-ws] Connecting...")
ws, _, err := wsDialer.Dial(b.connInfo.MediaServerURL, wsHeader)
if err != nil {
b.logger.Warn(fmt.Sprintf("[tm-ws] Connect failed: %s, retrying in 5s...", common.MaskError(err)))
time.Sleep(5 * time.Second)
continue
}
b.mu.Lock()
b.ws = ws
b.mu.Unlock()
b.logger.Debug("[tm-ws] Connected")
b.sendHello()
go b.pollAndAdmit()
stopWaitingRoomPoll := make(chan struct{})
go func() {
ticker := time.NewTicker(5 * time.Second)
defer ticker.Stop()
for {
select {
case <-stopWaitingRoomPoll:
return
case <-ticker.C:
b.pollAndAdmit()
}
}
}()
stopPing := make(chan struct{})
go func() {
ticker := time.NewTicker(tmPingPeriod)
defer ticker.Stop()
for {
select {
case <-stopPing:
return
case <-ticker.C:
b.wsSend(map[string]interface{}{"uid": uuid.New().String(), "ping": map[string]interface{}{}})
}
}
}()
stopStateKeepalive := make(chan struct{})
go func() {
interval := b.connInfo.StateCheckIntervalS
if interval <= 0 {
interval = 30
}
if err := b.requestStates(); err != nil {
b.logger.Debug(fmt.Sprintf("[tm-state] initial request-states: %v", err))
}
ticker := time.NewTicker(time.Duration(interval) * time.Second)
defer ticker.Stop()
for {
select {
case <-stopStateKeepalive:
return
case <-ticker.C:
if err := b.requestStates(); err != nil {
b.logger.Debug(fmt.Sprintf("[tm-state] request-states: %v", err))
}
}
}
}()
for {
_, raw, err := ws.ReadMessage()
if err != nil {
b.logger.Debug(fmt.Sprintf("[tm-ws] Closed: %s", common.MaskError(err)))
break
}
b.handleMessage(raw)
}
close(stopPing)
close(stopStateKeepalive)
close(stopWaitingRoomPoll)
b.mu.Lock()
b.ws = nil
b.mu.Unlock()
b.logger.Debug("[tm-ws] Rejoining in 3s...")
time.Sleep(3 * time.Second)
newConn, err := getConnection(b.dialer, b.cookieStr, url.QueryEscape(b.connInfo.ConferenceURI), b.config)
if err != nil {
b.logger.Warn(fmt.Sprintf("[rejoin] Failed: %v, retrying in 5s...", err))
time.Sleep(5 * time.Second)
continue
}
b.connInfo.PeerID = newConn.PeerID
b.connInfo.Credentials = newConn.Credentials
b.connInfo.MediaServerURL = newConn.MediaServerURL
b.connInfo.ICEServers = newConn.ICEServers
b.connInfo.StateCheckIntervalS = newConn.StateCheckIntervalS
}
}
func parseMids(sdp string) (audioMid, videoMid string) {
var media string
for _, line := range strings.Split(sdp, "\r\n") {
if strings.HasPrefix(line, "m=audio") {
media = "audio"
} else if strings.HasPrefix(line, "m=video") {
media = "video"
}
if strings.HasPrefix(line, "a=mid:") {
mid := strings.TrimPrefix(line, "a=mid:")
if media == "audio" && audioMid == "" {
audioMid = mid
} else if media == "video" && videoMid == "" {
videoMid = mid
}
}
}
return
}
func extractUfrag(candidate string) string {
parts := strings.Split(candidate, " ")
for i, p := range parts {
if p == "ufrag" && i+1 < len(parts) {
return parts[i+1]
}
}
return ""
}

View File

@@ -0,0 +1,933 @@
package telemost
import (
"context"
"crypto/tls"
"encoding/json"
"fmt"
"net"
"net/http"
"net/netip"
"net/url"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/dialer"
"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"
)
const (
TmAPIBase = APIBase
TmOrigin = Origin
TmPingPeriod = 5 * time.Second
telemostReconnectInitialDelay = time.Second
telemostReconnectMaxDelay = 16 * time.Second
)
type TelemostJoiner struct {
logger logger.ContextLogger
OnConnected func(tunnel.DataTunnel)
OnRemoteCandidate func(target int, candidateOrSDP string)
dialer N.Dialer
dnsRouter adapter.DNSRouter
PCConfig common.PeerConnectionConfigurer
AddTracks common.AddTunnelTracksFunc
ReadTrackFn common.ReadTrackFunc
joinLink string
displayName string
ws *websocket.Conn
wsMu sync.Mutex
subPC *webrtc.PeerConnection
subSeq int
subRemoteSet bool
subPending []webrtc.ICECandidateInit
pubPC *webrtc.PeerConnection
pubSeq int
pubRemoteSet bool
pubPending []webrtc.ICECandidateInit
sampleTrack *webrtc.TrackLocalStaticSample
vp8tunnel *tunnel.VP8DataTunnel
obf *tunnel.TunnelObfuscator
vp8FPS int
vp8Batch int
httpClient *http.Client
instanceID string
peerID string
roomID string
credentials string
serviceName string
mediaURL string
iceServers []webrtc.ICEServer
stateCheckIntervalS int
closeMu sync.Mutex
closed bool
stopCh chan struct{}
stopOnce sync.Once
configAck tunnel.ConfigAckTracker
reconnectAttempt atomic.Int32
setSlotsKey int
initBundleSent bool
boundPeers map[string]bool
unboundPeers map[string]bool
boundMu sync.Mutex
}
func NewTelemostJoiner(logger logger.ContextLogger, dialer N.Dialer, dnsRouter adapter.DNSRouter, pcConfig common.PeerConnectionConfigurer, addTracks common.AddTunnelTracksFunc, readTrackFn common.ReadTrackFunc) *TelemostJoiner {
return &TelemostJoiner{
logger: logger,
dialer: dialer,
dnsRouter: dnsRouter,
PCConfig: pcConfig,
AddTracks: addTracks,
ReadTrackFn: readTrackFn,
instanceID: uuid.New().String(),
stopCh: make(chan struct{}),
httpClient: &http.Client{
Timeout: 15 * time.Second,
Transport: &http.Transport{
TLSClientConfig: &tls.Config{InsecureSkipVerify: true},
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
return dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
},
},
},
}
}
func (j *TelemostJoiner) RunWithParams(jsonParams string) {
var params struct {
JoinLink string `json:"joinLink"`
DisplayName string `json:"displayName"`
VP8FPS int `json:"vp8Fps"`
VP8Batch int `json:"vp8Batch"`
}
if err := json.Unmarshal([]byte(jsonParams), &params); err != nil {
j.logger.Error(fmt.Sprintf("telemost-joiner: failed to parse params: %v", err))
return
}
j.joinLink = params.JoinLink
j.displayName = params.DisplayName
if j.displayName == "" {
j.displayName = "Joiner"
}
obf, err := tunnel.NewTunnelObfuscator(tunnel.DeriveSecretFromJoinLink(params.JoinLink))
if err != nil {
j.logger.Error(fmt.Sprintf("telemost-joiner: obfuscator init failed: %v", err))
return
}
j.obf = obf
j.vp8FPS = params.VP8FPS
j.vp8Batch = params.VP8Batch
j.logger.Info(fmt.Sprintf("telemost-joiner: link=%s name=%s vp8Fps=%d vp8Batch=%d localEpoch=0x%08x",
j.joinLink, j.displayName, params.VP8FPS, params.VP8Batch, obf.LocalEpoch()))
j.logger.Info("telemost-joiner: connecting")
if err := j.runOnce(); err != nil {
j.logger.Error(fmt.Sprintf("telemost-joiner: %v", err))
return
}
for {
if j.isClosed() {
return
}
j.logger.Info("telemost-joiner: tunnel lost")
j.resetSessionState()
if !j.waitBeforeRetry(int(j.reconnectAttempt.Load())) {
return
}
j.reconnectAttempt.Add(1)
if j.isClosed() {
return
}
j.logger.Info(fmt.Sprintf("telemost-joiner: reconnect attempt #%d", j.reconnectAttempt.Load()))
if err := j.runOnce(); err != nil {
j.logger.Warn(fmt.Sprintf("telemost-joiner: %v, will retry", err))
}
}
}
func (j *TelemostJoiner) Close() {
j.closeMu.Lock()
j.closed = true
j.closeMu.Unlock()
j.stopOnce.Do(func() { close(j.stopCh) })
j.wsMu.Lock()
ws := j.ws
j.ws = nil
j.wsMu.Unlock()
common.CloseWS(ws)
if j.vp8tunnel != nil {
j.vp8tunnel.Stop()
}
if j.subPC != nil {
j.subPC.Close()
}
if j.pubPC != nil {
j.pubPC.Close()
}
}
func TmParseMids(sdp string) (audioMid, videoMid string) {
var media string
for _, line := range strings.Split(sdp, "\r\n") {
if strings.HasPrefix(line, "m=audio") {
media = "audio"
} else if strings.HasPrefix(line, "m=video") {
media = "video"
}
if strings.HasPrefix(line, "a=mid:") {
mid := strings.TrimPrefix(line, "a=mid:")
if media == "audio" && audioMid == "" {
audioMid = mid
} else if media == "video" && videoMid == "" {
videoMid = mid
}
}
}
return
}
func (j *TelemostJoiner) runOnce() error {
if err := j.getConnection(); err != nil {
return err
}
j.connectAndRun()
return nil
}
func (j *TelemostJoiner) MarkConfigAcked() { j.configAck.Mark() }
func (j *TelemostJoiner) waitBeforeRetry(attempt int) bool {
delay := common.BackoffWithJitter(attempt, telemostReconnectInitialDelay, telemostReconnectMaxDelay)
j.logger.Debug(fmt.Sprintf("telemost-joiner: waiting %s before reconnect", delay))
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-timer.C:
return !j.isClosed()
case <-j.stopCh:
return false
}
}
func (j *TelemostJoiner) resetSessionState() {
j.wsMu.Lock()
j.ws = nil
j.wsMu.Unlock()
j.subPC = nil
j.subSeq = 0
j.subRemoteSet = false
j.subPending = nil
j.pubPC = nil
j.pubSeq = 0
j.pubRemoteSet = false
j.pubPending = nil
j.sampleTrack = nil
j.vp8tunnel = nil
j.initBundleSent = false
j.boundMu.Lock()
j.boundPeers = nil
j.unboundPeers = nil
j.boundMu.Unlock()
}
func (j *TelemostJoiner) isClosed() bool {
j.closeMu.Lock()
defer j.closeMu.Unlock()
return j.closed
}
func (j *TelemostJoiner) apiClient() *Client {
return &Client{HTTP: j.httpClient, InstanceID: j.instanceID}
}
func (j *TelemostJoiner) getConnection() error {
confURL := url.QueryEscape(j.joinLink)
name := url.QueryEscape(j.displayName)
if name == "" {
name = "Joiner"
}
connPath := "/conferences/" + confURL + "/connection?next_gen_media_platform_allowed=true&display_name=" + name + "&waiting_room_supported=true"
j.logger.Debug(fmt.Sprintf("telemost-joiner: getting connection for %s", j.joinLink))
responseBody, status, err := j.apiClient().TMRequest("GET", connPath)
if err != nil {
return fmt.Errorf("get connection: %w", err)
}
if status != 200 {
return fmt.Errorf("get connection: status %d: %s", status, string(responseBody))
}
var initial struct {
ConnectionType string `json:"connection_type"`
ClientConfig struct {
CheckInterval int `json:"conference_check_access_interval_ms"`
} `json:"client_configuration"`
}
json.Unmarshal(responseBody, &initial)
if initial.ConnectionType == "WAITING_ROOM" {
interval := initial.ClientConfig.CheckInterval
if interval <= 0 {
interval = 3000
}
checkPath := "/conferences/" + confURL + "/waiting-rooms/check-access"
j.logger.Info(fmt.Sprintf("telemost-joiner: in waiting room, polling check-access every %dms...", interval))
for {
time.Sleep(time.Duration(interval) * time.Millisecond)
checkBody, checkStatus, checkErr := j.apiClient().TMRequest("GET", checkPath)
if checkErr != nil {
return fmt.Errorf("waiting room check-access: %w", checkErr)
}
if checkStatus != 200 {
return fmt.Errorf("waiting room check-access: status %d", checkStatus)
}
var check struct {
Admitted bool `json:"admitted"`
}
json.Unmarshal(checkBody, &check)
if check.Admitted {
j.logger.Info("telemost-joiner: admitted!")
break
}
}
responseBody, status, err = j.apiClient().TMRequest("GET", connPath)
if err != nil {
return fmt.Errorf("post-admit connection: %w", err)
}
if status != 200 {
return fmt.Errorf("post-admit connection: status %d: %s", status, string(responseBody))
}
}
var conn struct {
PeerID string `json:"peer_id"`
RoomID string `json:"room_id"`
Credentials string `json:"credentials"`
ClientConfig struct {
MediaServerURL string `json:"media_server_url"`
ServiceName string `json:"service_name"`
ICEServers json.RawMessage `json:"ice_servers"`
StateCheckIntervalSecs int `json:"state_check_interval_seconds"`
} `json:"client_configuration"`
}
json.Unmarshal(responseBody, &conn)
if conn.ClientConfig.MediaServerURL == "" {
return fmt.Errorf("empty media_server_url: %s", string(responseBody))
}
j.peerID = conn.PeerID
j.roomID = conn.RoomID
j.credentials = conn.Credentials
j.mediaURL = conn.ClientConfig.MediaServerURL
j.serviceName = conn.ClientConfig.ServiceName
j.stateCheckIntervalS = conn.ClientConfig.StateCheckIntervalSecs
var rawIce []struct {
URLs []string `json:"urls"`
Username string `json:"username"`
Credential string `json:"credential"`
}
json.Unmarshal(conn.ClientConfig.ICEServers, &rawIce)
for _, s := range rawIce {
ice := webrtc.ICEServer{URLs: s.URLs}
if s.Username != "" {
ice.Username = s.Username
ice.Credential = s.Credential
}
j.iceServers = append(j.iceServers, ice)
}
j.logger.Debug(fmt.Sprintf("telemost-joiner: peer_id=%s room_id=%s media_url=%s", j.peerID, j.roomID, j.mediaURL))
return nil
}
func (j *TelemostJoiner) wsSend(msg interface{}) {
j.wsMu.Lock()
defer j.wsMu.Unlock()
if j.ws != nil {
data, _ := json.Marshal(msg)
j.logger.Debug(fmt.Sprintf("telemost-joiner: [DIAG] -> %s", string(data)))
j.ws.WriteJSON(msg)
}
}
func (j *TelemostJoiner) ack(uid string) {
if uid == "" {
return
}
j.wsSend(map[string]interface{}{
"uid": uid,
"ack": map[string]interface{}{
"status": map[string]interface{}{"code": "OK", "description": ""},
},
})
}
func (j *TelemostJoiner) sendHello() {
j.wsSend(map[string]interface{}{
"uid": uuid.New().String(),
"hello": map[string]interface{}{
"participantMeta": map[string]interface{}{"name": j.displayName, "role": "SPEAKER", "description": "", "sendAudio": false, "sendVideo": true},
"participantAttributes": map[string]interface{}{"name": j.displayName, "role": "SPEAKER", "description": ""},
"sendAudio": false, "sendVideo": true, "sendSharing": false,
"participantId": j.peerID,
"roomId": j.roomID,
"serviceName": j.serviceName,
"credentials": j.credentials,
"capabilitiesOffer": CapabilitiesOffer,
"sdkInfo": map[string]interface{}{"implementation": "browser", "version": "6.0.0", "userAgent": common.UserAgent, "hwConcurrency": 8},
"sdkInitializationId": uuid.New().String(),
"disablePublisher": false, "disableSubscriber": false, "disableSubscriberAudio": false,
},
})
j.logger.Debug("telemost-joiner: -> hello")
}
func (j *TelemostJoiner) sendICE(cand *webrtc.ICECandidate, target string, pcSeq int) {
candidate := cand.ToJSON()
j.wsSend(map[string]interface{}{
"uid": uuid.New().String(),
"webrtcIceCandidate": map[string]interface{}{
"candidate": candidate.Candidate, "sdpMid": *candidate.SDPMid,
"sdpMlineIndex": *candidate.SDPMLineIndex, "target": target, "pcSeq": pcSeq,
},
})
}
func (j *TelemostJoiner) initPC() {
config := webrtc.Configuration{ICEServers: j.iceServers}
settingEngine := webrtc.SettingEngine{}
settingEngine.DetachDataChannels()
if j.PCConfig != nil {
j.PCConfig.ConfigureSettingEngine(&settingEngine)
}
api, err := NewAPI(&settingEngine)
if err != nil {
j.logger.Error(fmt.Sprintf("telemost-joiner: ERROR: create webrtc API: %v", err))
return
}
subPC, err := api.NewPeerConnection(config)
if err != nil {
j.logger.Error(fmt.Sprintf("telemost-joiner: ERROR: create sub PC: %v", err))
return
}
j.subPC = subPC
subPC.OnICECandidate(func(cand *webrtc.ICECandidate) {
if cand != nil {
j.sendICE(cand, "SUBSCRIBER", j.subSeq)
}
})
subPC.OnConnectionStateChange(func(state webrtc.PeerConnectionState) {
j.logger.Debug(fmt.Sprintf("telemost-joiner: sub PC state: %s", state.String()))
if state == webrtc.PeerConnectionStateFailed {
j.logger.Error("telemost-joiner: ERROR: subscriber connection failed")
}
})
subPC.OnTrack(func(track *webrtc.TrackRemote, receiver *webrtc.RTPReceiver) {
j.logger.Debug(fmt.Sprintf("telemost-joiner: sub remote track: %s", track.Codec().MimeType))
go j.ReadTrackFn(track, func(frame []byte) {
if j.vp8tunnel != nil {
j.vp8tunnel.HandleFrame(frame)
}
}, j.logger, "telemost-joiner")
})
pubPC, err := api.NewPeerConnection(config)
if err != nil {
j.logger.Error(fmt.Sprintf("telemost-joiner: ERROR: create pub PC: %v", err))
return
}
j.pubPC = pubPC
j.pubSeq = 1
j.sampleTrack = j.AddTracks(pubPC, j.logger, "telemost-joiner [pub]")
pubPC.OnICECandidate(func(cand *webrtc.ICECandidate) {
if cand != nil {
j.sendICE(cand, "PUBLISHER", j.pubSeq)
}
})
pubPC.OnConnectionStateChange(func(state webrtc.PeerConnectionState) {
j.logger.Debug(fmt.Sprintf("telemost-joiner: pub PC state: %s", state.String()))
if state == webrtc.PeerConnectionStateConnected && j.vp8tunnel == nil {
j.reconnectAttempt.Store(0)
j.logger.Info("telemost-joiner: === VP8 TUNNEL CONNECTED ===")
j.vp8tunnel = tunnel.NewVP8DataTunnel(j.sampleTrack, j.obf, j.logger)
vp8tun := j.vp8tunnel
vp8tun.Start(j.vp8FPS, j.vp8Batch)
if !j.configAck.Acknowledged() {
acked, cancel := j.configAck.Arm()
go tunnel.SendVP8ConfigUntilAcked(acked, cancel, j.stopCh, vp8tun,
vp8tun.FPS(), vp8tun.Batch(), 1, j.logger, "telemost-joiner")
j.logger.Debug(fmt.Sprintf("telemost-joiner: pushed vp8 config to creator fps=%d batch=%d", vp8tun.FPS(), vp8tun.Batch()))
}
if j.OnConnected != nil {
j.OnConnected(j.vp8tunnel)
}
}
})
j.logger.Debug(fmt.Sprintf("telemost-joiner: sub+pub PCs created with %d ICE servers", len(j.iceServers)))
}
func (j *TelemostJoiner) sendPubOffer() {
if j.pubPC == nil {
return
}
offer, err := j.pubPC.CreateOffer(nil)
if err != nil {
j.logger.Warn(fmt.Sprintf("telemost-joiner: pub offer failed: %v", err))
return
}
if err := j.pubPC.SetLocalDescription(offer); err != nil {
j.logger.Warn(fmt.Sprintf("telemost-joiner: set pub local desc: %v", err))
return
}
offer.SDP = MungeSDPAddVideoContent(offer.SDP)
audioMid, videoMid := TmParseMids(offer.SDP)
j.logger.Debug(fmt.Sprintf("telemost-joiner: -> publisherSdpOffer pcSeq=%d audioMid=%s videoMid=%s", j.pubSeq, audioMid, videoMid))
var tracks []map[string]interface{}
if audioMid != "" {
tracks = append(tracks, map[string]interface{}{"mid": audioMid, "transceiverMid": audioMid, "kind": "AUDIO", "priority": 0, "label": "", "codecs": map[string]interface{}{}, "groupId": 1, "description": ""})
}
if videoMid != "" {
tracks = append(tracks, map[string]interface{}{"mid": videoMid, "transceiverMid": videoMid, "kind": "VIDEO", "priority": 0, "label": "", "codecs": map[string]interface{}{}, "groupId": 2, "description": ""})
}
j.wsSend(map[string]interface{}{
"uid": uuid.New().String(),
"publisherSdpOffer": map[string]interface{}{"pcSeq": j.pubSeq, "sdp": offer.SDP, "tracks": tracks},
})
}
func (j *TelemostJoiner) handlePubAnswer(sdp string) {
if j.pubPC == nil {
return
}
if j.OnRemoteCandidate != nil {
j.OnRemoteCandidate(-1, sdp)
}
err := j.pubPC.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeAnswer,
SDP: sdp,
})
if err != nil {
j.logger.Warn(fmt.Sprintf("telemost-joiner: set pub remote desc: %v", err))
return
}
j.pubRemoteSet = true
for _, candidate := range j.pubPending {
j.pubPC.AddICECandidate(candidate)
}
j.pubPending = nil
j.sendInitBundle()
}
func (j *TelemostJoiner) sendInitBundle() {
if j.initBundleSent {
return
}
j.initBundleSent = true
j.logger.Debug("telemost-joiner: -> sdkCodecsInfo + updatePublisherTrackDescription")
j.wsSend(SdkCodecsInfoMessage())
j.wsSend(UpdatePublisherTrackDescriptionMessage(j.pubPC, "Microphone", "MacBook Pro Camera (0000:0001)"))
j.sendStartupSlotsRamp()
}
func (j *TelemostJoiner) requestVideoSlots() {
j.setSlotsKey++
j.logger.Debug(fmt.Sprintf("telemost-joiner: -> setSlots key=%d", j.setSlotsKey))
j.wsSend(SetSlotsMessage(j.setSlotsKey))
}
func (j *TelemostJoiner) forceReconnect(reason string) {
j.reconnectAttempt.Store(0)
oldPeerID := j.peerID
j.logger.Info(fmt.Sprintf("telemost-joiner: forcing reconnect: %s", reason))
if oldPeerID != "" {
j.logger.Debug(fmt.Sprintf("telemost-joiner: kicking self pid=%s to leave call cleanly", oldPeerID))
confURL := url.QueryEscape(j.joinLink)
_, status, err := j.apiClient().TMRequest("POST", "/conferences/"+confURL+"/commands/kick?peer_id="+url.QueryEscape(oldPeerID)+"&with_ban=false")
if err != nil || status >= 400 {
j.logger.Warn(fmt.Sprintf("telemost-joiner: self-kick failed: status=%d err=%v", status, err))
}
}
j.instanceID = uuid.New().String()
j.logger.Debug(fmt.Sprintf("telemost-joiner: new instance-id=%s", j.instanceID))
j.wsMu.Lock()
ws := j.ws
j.wsMu.Unlock()
common.CloseWS(ws)
}
func (j *TelemostJoiner) sendStartupSlotsRamp() {
for i := 0; i < 4; i++ {
j.setSlotsKey++
j.logger.Debug(fmt.Sprintf("telemost-joiner: -> setSlots key=%d (startup %d/4)", j.setSlotsKey, i+1))
j.wsSend(StartupSetSlotsMessage(i, j.setSlotsKey))
}
}
func (j *TelemostJoiner) handleSubOffer(sdp string, pcSeq int) {
j.subSeq = pcSeq
if j.subPC == nil {
j.logger.Warn("telemost-joiner: sub PC not ready for offer")
return
}
if j.OnRemoteCandidate != nil {
j.OnRemoteCandidate(-1, sdp)
}
err := j.subPC.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeOffer,
SDP: sdp,
})
if err != nil {
j.logger.Warn(fmt.Sprintf("telemost-joiner: set sub remote desc: %v", err))
return
}
j.subRemoteSet = true
for _, candidate := range j.subPending {
j.subPC.AddICECandidate(candidate)
}
j.subPending = nil
answer, err := j.subPC.CreateAnswer(nil)
if err != nil {
j.logger.Warn(fmt.Sprintf("telemost-joiner: create sub answer: %v", err))
return
}
j.subPC.SetLocalDescription(answer)
j.logger.Debug(fmt.Sprintf("telemost-joiner: -> subscriberSdpAnswer pcSeq=%d", pcSeq))
j.wsSend(map[string]interface{}{
"uid": uuid.New().String(),
"subscriberSdpAnswer": map[string]interface{}{"sdp": answer.SDP, "pcSeq": pcSeq},
})
j.sendPubOffer()
}
func (j *TelemostJoiner) handleMessage(raw []byte) {
var msg map[string]interface{}
if err := json.Unmarshal(raw, &msg); err != nil {
return
}
uid, _ := msg["uid"].(string)
if _, ok := msg["serverHello"]; ok {
j.logger.Debug("telemost-joiner: <- serverHello")
if sh, ok := msg["serverHello"].(map[string]interface{}); ok {
j.parseICEServersFromHello(sh)
}
j.ack(uid)
j.initPC()
return
}
if so, ok := msg["subscriberSdpOffer"]; ok {
soMap, _ := so.(map[string]interface{})
sdp, _ := soMap["sdp"].(string)
pcSeq, _ := soMap["pcSeq"].(float64)
j.logger.Debug(fmt.Sprintf("telemost-joiner: <- subscriberSdpOffer pcSeq=%d len=%d", int(pcSeq), len(sdp)))
j.ack(uid)
j.handleSubOffer(sdp, int(pcSeq))
return
}
if pa, ok := msg["publisherSdpAnswer"]; ok {
paMap, _ := pa.(map[string]interface{})
sdp, _ := paMap["sdp"].(string)
j.logger.Debug(fmt.Sprintf("telemost-joiner: <- publisherSdpAnswer %d bytes", len(sdp)))
j.handlePubAnswer(sdp)
return
}
if ic, ok := msg["webrtcIceCandidate"]; ok {
icMap, _ := ic.(map[string]interface{})
candidate, _ := icMap["candidate"].(string)
sdpMid, _ := icMap["sdpMid"].(string)
target, _ := icMap["target"].(string)
sdpIdx, _ := icMap["sdpMlineIndex"].(float64)
idx := uint16(sdpIdx)
cand := webrtc.ICECandidateInit{Candidate: candidate, SDPMid: &sdpMid, SDPMLineIndex: &idx}
if j.OnRemoteCandidate != nil {
tgt := 1
if target == "SUBSCRIBER" {
tgt = 0
}
j.OnRemoteCandidate(tgt, candidate)
}
if target == "SUBSCRIBER" {
if j.subRemoteSet {
j.subPC.AddICECandidate(cand)
} else {
j.subPending = append(j.subPending, cand)
}
} else if target == "PUBLISHER" {
if j.pubRemoteSet {
j.pubPC.AddICECandidate(cand)
} else {
j.pubPending = append(j.pubPending, cand)
}
}
j.ack(uid)
return
}
if ackData, ok := msg["ack"]; ok {
if ackMap, ok := ackData.(map[string]interface{}); ok {
if status, ok := ackMap["status"].(map[string]interface{}); ok {
if code, _ := status["code"].(string); code != "OK" {
desc, _ := status["description"].(string)
j.logger.Warn(fmt.Sprintf("telemost-joiner: ack error: %s %s", code, desc))
}
}
}
return
}
if ud, ok := msg["upsertDescription"]; ok {
udMap, _ := ud.(map[string]interface{})
if descs, ok := udMap["description"].([]interface{}); ok {
for _, d := range descs {
dm, _ := d.(map[string]interface{})
pid, _ := dm["id"].(string)
if pid != "" && pid != j.peerID {
participantName := ""
if meta, ok := dm["meta"].(map[string]interface{}); ok {
participantName, _ = meta["name"].(string)
}
j.logger.Debug(fmt.Sprintf("telemost-joiner: participant: %s (%s)", participantName, pid))
}
}
}
j.ack(uid)
return
}
if ud, ok := msg["updateDescription"]; ok {
j.logger.Debug(fmt.Sprintf("telemost-joiner: <- updateDescription %s", BriefJSON(ud)))
j.ack(uid)
return
}
if _, ok := msg["removeDescription"]; ok {
j.logger.Info("telemost-joiner: participant left")
j.ack(uid)
return
}
if sc, ok := msg["slotsConfig"]; ok {
j.logger.Debug(fmt.Sprintf("telemost-joiner: <- slotsConfig %s", BriefJSON(sc)))
needRebind := false
presentPids := make(map[string]bool)
for _, ev := range SlotsConfigBindings(sc) {
fullPid := ev.ParticipantID
if fullPid != "" {
presentPids[fullPid] = true
}
pid := fullPid
if len(pid) > 8 {
pid = pid[:8]
}
if ev.Reason == "NO_LIMITATION" && ev.Mid != "" {
j.logger.Debug(fmt.Sprintf("telemost-joiner: [bind] BOUND slot=%d pid=%s mid=%s", ev.Slot, pid, ev.Mid))
j.boundMu.Lock()
if j.boundPeers == nil {
j.boundPeers = make(map[string]bool)
}
j.boundPeers[fullPid] = true
delete(j.unboundPeers, fullPid)
j.boundMu.Unlock()
} else if fullPid != "" {
j.boundMu.Lock()
wasBound := j.boundPeers[fullPid]
if wasBound {
if j.unboundPeers == nil {
j.unboundPeers = make(map[string]bool)
}
j.unboundPeers[fullPid] = true
delete(j.boundPeers, fullPid)
}
j.boundMu.Unlock()
if wasBound {
j.logger.Debug(fmt.Sprintf("telemost-joiner: [bind] KILL slot=%d pid=%s reason=%s - rebinding", ev.Slot, pid, ev.Reason))
needRebind = true
} else {
j.logger.Debug(fmt.Sprintf("telemost-joiner: [bind] UNBOUND slot=%d pid=%s reason=%s mid=%q", ev.Slot, pid, ev.Reason, ev.Mid))
}
}
}
j.boundMu.Lock()
for boundPid := range j.boundPeers {
if !presentPids[boundPid] {
short := boundPid
if len(short) > 8 {
short = short[:8]
}
j.logger.Debug(fmt.Sprintf("telemost-joiner: [bind] VANISHED pid=%s - rebinding", short))
delete(j.boundPeers, boundPid)
needRebind = true
}
}
j.boundMu.Unlock()
if needRebind {
go j.forceReconnect("slot binding killed")
}
j.ack(uid)
return
}
for k, v := range msg {
if k == "uid" || k == "ack" {
continue
}
j.logger.Debug(fmt.Sprintf("telemost-joiner: <- %s (unhandled) %s", k, BriefJSON(v)))
break
}
if uid != "" {
j.ack(uid)
}
}
func (j *TelemostJoiner) parseICEServersFromHello(sh map[string]interface{}) {
rtcCfg, ok := sh["rtcConfiguration"].(map[string]interface{})
if !ok {
return
}
servers, ok := rtcCfg["iceServers"].([]interface{})
if !ok {
return
}
var iceServers []webrtc.ICEServer
for _, s := range servers {
sm, _ := s.(map[string]interface{})
var urls []string
if u, ok := sm["urls"].([]interface{}); ok {
for _, v := range u {
if vs, ok := v.(string); ok {
urls = append(urls, common.FixICEURL(vs))
}
}
}
ice := webrtc.ICEServer{URLs: urls}
if u, ok := sm["username"].(string); ok && u != "" {
ice.Username = u
ice.Credential, _ = sm["credential"].(string)
}
iceServers = append(iceServers, ice)
}
resolved := make(map[string]string)
for i, s := range iceServers {
for k, u := range s.URLs {
host := common.ExtractICEHost(u)
if host == "" || net.ParseIP(host) != nil {
continue
}
_, ok := resolved[host]
if !ok {
rd, hasRD := j.dialer.(dialer.ResolveDialer)
if j.dnsRouter == nil || !hasRD {
continue
}
var err error
var addrs []netip.Addr
addrs, err = j.dnsRouter.Lookup(context.Background(), host, rd.QueryOptions())
if err != nil {
j.logger.Warn(fmt.Sprintf("telemost-joiner: resolve ICE host %s failed: %s", common.MaskAddr(host), common.MaskError(err)))
continue
}
resolved[host] = addrs[0].String()
j.logger.Debug(fmt.Sprintf("telemost-joiner: resolved ICE host %s -> %s", host, addrs[0]))
}
iceServers[i].URLs[k] = strings.Replace(u, host, resolved[host], 1)
}
}
j.iceServers = iceServers
for i, s := range iceServers {
j.logger.Debug(fmt.Sprintf("telemost-joiner: ICE server %d: urls=%v", i, s.URLs))
}
j.logger.Debug(fmt.Sprintf("telemost-joiner: %d ICE servers from serverHello", len(iceServers)))
}
func (j *TelemostJoiner) connectAndRun() {
parsed, err := url.Parse(j.mediaURL)
if err != nil {
j.logger.Error(fmt.Sprintf("telemost-joiner: ERROR: bad media URL: %s", common.MaskError(err)))
return
}
hostname := parsed.Hostname()
wsHeader := http.Header{}
wsHeader.Set("User-Agent", common.UserAgent)
wsHeader.Set("Origin", TmOrigin)
j.logger.Debug(fmt.Sprintf("telemost-joiner: connecting to %s", j.mediaURL))
dialer := websocket.Dialer{
HandshakeTimeout: 10 * time.Second,
WriteBufferSize: 65536,
TLSClientConfig: &tls.Config{InsecureSkipVerify: true, ServerName: hostname},
NetDialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
return j.dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
},
}
ws, _, err := dialer.Dial(j.mediaURL, wsHeader)
if err != nil {
j.logger.Error(fmt.Sprintf("telemost-joiner: ERROR: ws connect: %s", common.MaskError(err)))
return
}
j.wsMu.Lock()
j.ws = ws
j.wsMu.Unlock()
j.logger.Debug("telemost-joiner: ws connected")
j.sendHello()
stopPing := make(chan struct{})
go func() {
ticker := time.NewTicker(TmPingPeriod)
defer ticker.Stop()
for {
select {
case <-stopPing:
return
case <-ticker.C:
j.wsSend(map[string]interface{}{"uid": uuid.New().String(), "ping": map[string]interface{}{}})
}
}
}()
stopStateKeepalive := make(chan struct{})
go func() {
interval := j.stateCheckIntervalS
if interval <= 0 {
interval = 30
}
if err := j.apiClient().RequestStates(j.joinLink, j.peerID); err != nil {
j.logger.Debug(fmt.Sprintf("telemost-joiner: initial request-states: %v", err))
}
ticker := time.NewTicker(time.Duration(interval) * time.Second)
defer ticker.Stop()
for {
select {
case <-stopStateKeepalive:
return
case <-ticker.C:
if err := j.apiClient().RequestStates(j.joinLink, j.peerID); err != nil {
j.logger.Debug(fmt.Sprintf("telemost-joiner: request-states: %v", err))
}
}
}
}()
for {
_, raw, err := ws.ReadMessage()
if err != nil {
j.logger.Debug(fmt.Sprintf("telemost-joiner: ws read error: %s", common.MaskError(err)))
break
}
j.handleMessage(raw)
}
close(stopPing)
close(stopStateKeepalive)
if j.vp8tunnel != nil {
j.vp8tunnel.Stop()
}
if j.subPC != nil {
j.subPC.Close()
}
if j.pubPC != nil {
j.pubPC.Close()
}
j.logger.Info("telemost-joiner: disconnected")
}

View File

@@ -0,0 +1,267 @@
package telemost
import (
"fmt"
"sync"
"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"
)
type SFURelay struct {
pubPC *webrtc.PeerConnection
subPC *webrtc.PeerConnection
pubRemoteSet bool
subRemoteSet bool
pubPending []webrtc.ICECandidateInit
subPending []webrtc.ICECandidateInit
mu sync.Mutex
logger logger.ContextLogger
sampleTrack *webrtc.TrackLocalStaticSample
tun *tunnel.VP8DataTunnel
obf *tunnel.TunnelObfuscator
OnConnected func(*tunnel.VP8DataTunnel)
OnPubReady func()
OnPeerRestart func()
OnPubICE func(*webrtc.ICECandidate)
OnSubICE func(*webrtc.ICECandidate)
readBufSize int
}
func (r *SFURelay) SetObfuscator(o *tunnel.TunnelObfuscator) { r.obf = o }
func NewSFURelay(logger logger.ContextLogger) *SFURelay {
return &SFURelay{logger: logger}
}
func (r *SFURelay) Init(iceServers []webrtc.ICEServer) error {
config := webrtc.Configuration{ICEServers: iceServers}
pubPC, err := NewPeerConnection(config)
if err != nil {
return err
}
r.pubPC = pubPC
sampleTrack, _ := webrtc.NewTrackLocalStaticSample(
webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeVP8},
"video", "tunnel-video",
)
r.sampleTrack = sampleTrack
audioTrack, _ := webrtc.NewTrackLocalStaticRTP(
webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeOpus},
"audio", "tunnel-audio",
)
pubPC.AddTransceiverFromTrack(audioTrack, webrtc.RTPTransceiverInit{Direction: webrtc.RTPTransceiverDirectionSendonly})
pubPC.AddTransceiverFromTrack(r.sampleTrack, webrtc.RTPTransceiverInit{Direction: webrtc.RTPTransceiverDirectionSendonly})
pubPC.OnICECandidate(func(cand *webrtc.ICECandidate) {
if cand == nil || r.OnPubICE == nil {
return
}
r.OnPubICE(cand)
})
pubPC.OnConnectionStateChange(func(state webrtc.PeerConnectionState) {
r.logger.Debug(fmt.Sprintf("[pub] connection state: %s", state.String()))
if state == webrtc.PeerConnectionStateConnected {
if r.tun == nil {
r.logger.Debug("[relay] starting VP8 publish tunnel on pub PC connected")
r.tun = tunnel.NewVP8DataTunnel(r.sampleTrack, r.obf, r.logger)
r.tun.Start(0, 0)
if r.OnConnected != nil {
r.OnConnected(r.tun)
}
}
if r.OnPubReady != nil {
r.OnPubReady()
}
}
})
subPC, err := NewPeerConnection(config)
if err != nil {
pubPC.Close()
return err
}
r.subPC = subPC
subPC.OnICECandidate(func(cand *webrtc.ICECandidate) {
if cand == nil || r.OnSubICE == nil {
return
}
r.OnSubICE(cand)
})
subPC.OnConnectionStateChange(func(state webrtc.PeerConnectionState) {
r.logger.Debug(fmt.Sprintf("[sub] connection state: %s", state.String()))
})
subPC.OnTrack(func(track *webrtc.TrackRemote, receiver *webrtc.RTPReceiver) {
r.logger.Debug(fmt.Sprintf("[sub] remote track: %s", track.Codec().MimeType))
go r.readTrack(track)
})
r.logger.Debug(fmt.Sprintf("[relay] pub+sub PCs created (%d ICE servers)", len(iceServers)))
return nil
}
func (r *SFURelay) CreatePubOffer() (webrtc.SessionDescription, error) {
offer, err := r.pubPC.CreateOffer(nil)
if err != nil {
return offer, err
}
if err := r.pubPC.SetLocalDescription(offer); err != nil {
return offer, err
}
offer.SDP = MungeSDPAddVideoContent(offer.SDP)
return offer, nil
}
func (r *SFURelay) SetPubAnswer(sdp string) error {
r.mu.Lock()
defer r.mu.Unlock()
err := r.pubPC.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeAnswer, SDP: sdp,
})
if err != nil {
return err
}
r.pubRemoteSet = true
for _, cand := range r.pubPending {
r.pubPC.AddICECandidate(cand)
}
r.pubPending = nil
return nil
}
func (r *SFURelay) SetSubOffer(sdp string) (webrtc.SessionDescription, error) {
r.mu.Lock()
defer r.mu.Unlock()
err := r.subPC.SetRemoteDescription(webrtc.SessionDescription{
Type: webrtc.SDPTypeOffer, SDP: sdp,
})
if err != nil {
return webrtc.SessionDescription{}, err
}
r.subRemoteSet = true
for _, cand := range r.subPending {
r.subPC.AddICECandidate(cand)
}
r.subPending = nil
answer, err := r.subPC.CreateAnswer(nil)
if err != nil {
return answer, err
}
r.subPC.SetLocalDescription(answer)
return answer, nil
}
func (r *SFURelay) AddPubICECandidate(cand webrtc.ICECandidateInit) {
r.mu.Lock()
defer r.mu.Unlock()
if !r.pubRemoteSet {
r.pubPending = append(r.pubPending, cand)
return
}
r.pubPC.AddICECandidate(cand)
}
func (r *SFURelay) AddSubICECandidate(cand webrtc.ICECandidateInit) {
r.mu.Lock()
defer r.mu.Unlock()
if !r.subRemoteSet {
r.subPending = append(r.subPending, cand)
return
}
r.subPC.AddICECandidate(cand)
}
func (r *SFURelay) Close() {
if r.tun != nil {
r.tun.Stop()
r.tun = nil
}
if r.pubPC != nil {
r.pubPC.Close()
r.pubPC = nil
}
if r.subPC != nil {
r.subPC.Close()
r.subPC = nil
}
}
func (r *SFURelay) 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
bufSz := r.readBufSize
if bufSz <= 0 {
bufSz = common.RTPBufSize
}
buf := make([]byte, bufSz)
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 {
r.logger.Debug(fmt.Sprintf("[video] recv vp8 frame #%d %d bytes", recvCount, len(frameBuf)))
}
res := r.obf.Decode(frameBuf)
frameBuf = frameBuf[:0]
frameValid = false
if !res.HasFrame || res.SelfEcho {
continue
}
if res.PeerRestart {
r.logger.Info(fmt.Sprintf("[video] peer restart detected, new epoch=0x%08x", res.PeerEpoch))
if r.OnPeerRestart != nil {
r.OnPeerRestart()
}
}
if res.Keepalive || len(res.Payload) == 0 {
continue
}
if r.tun != nil && r.tun.OnData != nil {
r.tun.OnData(res.Payload)
}
}
}

View File

@@ -0,0 +1,69 @@
package tunnel
import (
"fmt"
"sync"
"time"
"github.com/sagernet/sing/common/logger"
)
const configResendPeriod = 3 * time.Second
type ConfigAckTracker struct {
mu sync.Mutex
acked chan struct{}
cancel chan struct{}
confirmed bool
}
func (t *ConfigAckTracker) Acknowledged() bool {
t.mu.Lock()
defer t.mu.Unlock()
return t.confirmed
}
func (t *ConfigAckTracker) Arm() (acked, cancel chan struct{}) {
t.mu.Lock()
defer t.mu.Unlock()
if t.cancel != nil {
close(t.cancel)
}
t.acked = make(chan struct{})
t.cancel = make(chan struct{})
return t.acked, t.cancel
}
func (t *ConfigAckTracker) Mark() {
t.mu.Lock()
defer t.mu.Unlock()
t.confirmed = true
if t.acked == nil {
return
}
select {
case <-t.acked:
default:
close(t.acked)
}
}
func SendVP8ConfigUntilAcked(acked, cancel <-chan struct{}, stopCh <-chan struct{}, tun DataTunnel, fps, batch, trackCount int, logger logger.ContextLogger, logPrefix string) {
tun.SendData(EncodeVP8Config(fps, batch, trackCount))
ticker := time.NewTicker(configResendPeriod)
defer ticker.Stop()
for {
select {
case <-acked:
return
case <-cancel:
return
case <-stopCh:
return
case <-ticker.C:
logger.Debug(fmt.Sprintf("%s: resending vp8 config fps=%d batch=%d trackCount=%d, no ack yet",
logPrefix, fps, batch, trackCount))
tun.SendData(EncodeVP8Config(fps, batch, trackCount))
}
}
}

View File

@@ -0,0 +1,230 @@
package tunnel
import (
"encoding/binary"
"fmt"
"io"
"math"
"sync"
"sync/atomic"
"github.com/pion/datachannel"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing/common/logger"
)
const chunkSize = 994
type chunkBuf struct {
chunks [][]byte
count int
size int
}
type DCTunnel struct {
dc *webrtc.DataChannel
raw datachannel.ReadWriteCloser
writeRaw datachannel.ReadWriteCloser
logger logger.ContextLogger
onData func([]byte)
onClose func()
obf *TunnelObfuscator
chunked bool
readBuf int
recvBufs sync.Map
sendMsgID uint32
}
func NewDCTunnel(dc *webrtc.DataChannel, obf *TunnelObfuscator, readBuf int, logger logger.ContextLogger) *DCTunnel {
t := &DCTunnel{dc: dc, obf: obf, readBuf: readBuf, logger: logger}
raw, err := dc.Detach()
if err != nil {
logger.Warn(fmt.Sprintf("dctunnel: detach failed, using callback mode: %v", err))
dc.OnMessage(func(msg webrtc.DataChannelMessage) {
t.deliverMessage(msg.Data)
})
dc.OnClose(func() {
if t.onClose != nil {
t.onClose()
}
})
return t
}
t.raw = raw
go t.readLoop()
return t
}
func NewDCTunnelFromRaw(dc *webrtc.DataChannel, raw datachannel.ReadWriteCloser, obf *TunnelObfuscator, readBuf int, logger logger.ContextLogger) *DCTunnel {
t := &DCTunnel{dc: dc, raw: raw, obf: obf, readBuf: readBuf, logger: logger}
go t.readLoop()
return t
}
func NewChunkedDCTunnel(readRaw datachannel.ReadWriteCloser, writeDC *webrtc.DataChannel, obf *TunnelObfuscator, readBuf int, logger logger.ContextLogger) *DCTunnel {
writeRaw, err := writeDC.Detach()
if err != nil {
logger.Error(fmt.Sprintf("dctunnel: write DC detach failed: %v", err))
return nil
}
t := &DCTunnel{raw: readRaw, writeRaw: writeRaw, obf: obf, readBuf: readBuf, logger: logger, chunked: true}
go t.readLoop()
return t
}
func NewChunkedDCTunnelFromRaw(readRaw, writeRaw datachannel.ReadWriteCloser, obf *TunnelObfuscator, readBuf int, logger logger.ContextLogger) *DCTunnel {
t := &DCTunnel{raw: readRaw, writeRaw: writeRaw, obf: obf, readBuf: readBuf, logger: logger, chunked: true}
go t.readLoop()
return t
}
func (t *DCTunnel) SendData(data []byte) {
for len(data) >= 4 {
frameLen := int(binary.BigEndian.Uint32(data[0:4]))
if frameLen < 5 || 4+frameLen > len(data) {
return
}
body := data[4 : 4+frameLen]
wire := body
if t.obf != nil {
wire = t.obf.EncryptPayload(body)
if wire == nil {
data = data[4+frameLen:]
continue
}
}
if t.chunked {
t.sendChunked(wire)
} else {
t.sendRaw(wire)
}
data = data[4+frameLen:]
}
}
func (t *DCTunnel) SetOnData(fn func([]byte)) { t.onData = fn }
func (t *DCTunnel) OnData() func([]byte) { return t.onData }
func (t *DCTunnel) SetOnClose(fn func()) { t.onClose = fn }
func (t *DCTunnel) Reconfigure(fps, batch int) {}
func (t *DCTunnel) readLoop() {
buf := make([]byte, t.readBuf)
for {
n, isString, err := t.raw.ReadDataChannel(buf)
if err != nil {
if err != io.EOF {
t.logger.Warn(fmt.Sprintf("dctunnel: read error: %v", err))
}
if t.onClose != nil {
t.onClose()
}
return
}
if isString {
continue
}
if t.chunked && n >= 6 {
t.handleChunk(buf[:n])
} else if n > 0 {
t.deliverMessage(buf[:n])
}
}
}
func (t *DCTunnel) handleChunk(data []byte) {
id := uint16(data[0])<<8 | uint16(data[1])
idx := int(uint16(data[2])<<8 | uint16(data[3]))
total := int(uint16(data[4])<<8 | uint16(data[5]))
payload := data[6:]
if total == 1 {
cp := make([]byte, len(payload))
copy(cp, payload)
t.deliverMessage(cp)
return
}
val, _ := t.recvBufs.LoadOrStore(id, &chunkBuf{chunks: make([][]byte, total)})
cb := val.(*chunkBuf)
if idx < len(cb.chunks) && cb.chunks[idx] == nil {
cp := make([]byte, len(payload))
copy(cp, payload)
cb.chunks[idx] = cp
cb.count++
cb.size += len(cp)
}
if cb.count == total {
t.recvBufs.Delete(id)
out := make([]byte, 0, cb.size)
for _, c := range cb.chunks {
out = append(out, c...)
}
t.deliverMessage(out)
}
}
func (t *DCTunnel) deliverMessage(data []byte) {
if len(data) == 0 {
return
}
if t.obf != nil {
pt, ok := t.obf.DecryptPayload(data)
if !ok {
t.logger.Debug(fmt.Sprintf("dctunnel: decrypt failed, dropping %d bytes", len(data)))
return
}
data = pt
}
if t.onData != nil && len(data) > 0 {
frame := make([]byte, 4+len(data))
binary.BigEndian.PutUint32(frame[0:4], uint32(len(data)))
copy(frame[4:], data)
t.onData(frame)
}
}
func (t *DCTunnel) sendChunked(data []byte) {
w := t.writeRaw
if w == nil {
w = t.raw
}
if w == nil {
return
}
total := int(math.Ceil(float64(len(data)) / float64(chunkSize)))
if total == 0 {
total = 1
}
id := uint16(atomic.AddUint32(&t.sendMsgID, 1)) & 0xFFFF
for i := 0; i < total; i++ {
start := i * chunkSize
end := start + chunkSize
if end > len(data) {
end = len(data)
}
p := data[start:end]
f := make([]byte, 6+len(p))
f[0] = byte(id >> 8)
f[1] = byte(id & 0xFF)
f[2] = byte(i >> 8)
f[3] = byte(i & 0xFF)
f[4] = byte(total >> 8)
f[5] = byte(total & 0xFF)
copy(f[6:], p)
w.Write(f)
}
}
func (t *DCTunnel) sendRaw(data []byte) {
w := t.writeRaw
if w == nil {
w = t.raw
}
if w != nil {
w.Write(data)
return
}
if t.dc == nil || t.dc.ReadyState() != webrtc.DataChannelStateOpen {
return
}
t.dc.Send(data)
}

View File

@@ -0,0 +1,69 @@
package tunnel
import (
"context"
"io"
"testing"
)
type discardRawConn struct{}
func (discardRawConn) Read(p []byte) (int, error) { return 0, io.EOF }
func (discardRawConn) ReadDataChannel(p []byte) (int, bool, error) { return 0, false, io.EOF }
func (discardRawConn) Write(p []byte) (int, error) { return len(p), nil }
func (discardRawConn) WriteDataChannel(p []byte, isString bool) (int, error) {
return len(p), nil
}
func (discardRawConn) Close() error { return nil }
type benchLogger struct{}
func (benchLogger) Trace(args ...any) {}
func (benchLogger) Debug(args ...any) {}
func (benchLogger) Info(args ...any) {}
func (benchLogger) Notice(args ...any) {}
func (benchLogger) Warn(args ...any) {}
func (benchLogger) Error(args ...any) {}
func (benchLogger) Fatal(args ...any) {}
func (benchLogger) Panic(args ...any) {}
func (benchLogger) TraceContext(ctx context.Context, args ...any) {}
func (benchLogger) DebugContext(ctx context.Context, args ...any) {}
func (benchLogger) InfoContext(ctx context.Context, args ...any) {}
func (benchLogger) NoticeContext(ctx context.Context, args ...any) {}
func (benchLogger) WarnContext(ctx context.Context, args ...any) {}
func (benchLogger) ErrorContext(ctx context.Context, args ...any) {}
func (benchLogger) FatalContext(ctx context.Context, args ...any) {}
func (benchLogger) PanicContext(ctx context.Context, args ...any) {}
func newBenchDCTunnel() *DCTunnel {
return &DCTunnel{raw: discardRawConn{}, logger: benchLogger{}, readBuf: 4096}
}
func BenchmarkDCTunnelSendData(b *testing.B) {
sizes := []int{64, 512, 4096}
for _, size := range sizes {
payload := make([]byte, size)
frame := EncodeFrame(42, MsgData, payload)
b.Run(sizeLabel(size), func(b *testing.B) {
t := newBenchDCTunnel()
b.ReportAllocs()
b.SetBytes(int64(len(frame)))
for i := 0; i < b.N; i++ {
t.SendData(frame)
}
})
}
}
func sizeLabel(n int) string {
switch n {
case 64:
return "64B"
case 512:
return "512B"
case 4096:
return "4KB"
default:
return "custom"
}
}

View File

@@ -0,0 +1,400 @@
package tunnel
import (
"encoding/binary"
"fmt"
"sync"
"sync/atomic"
"time"
kcp "github.com/xtaci/kcp-go/v5"
"github.com/sagernet/sing/common/logger"
)
const (
kcpConvBase = 0x77627374
kcpUpdateInterval = 10 * time.Millisecond
// One KCP segment must ride in a single RTP packet so a dropped packet
// loses only its own frame, not a two-packet frame that readVP8Track
// would discard whole. 1200 RTP budget - 1 VP8 descriptor - interframe
// header - 24 XChaCha20 nonce - 16 Poly1305 tag - 1 channel tag.
kcpSegmentMTU = 1200 - 1 - interframeHdrLen - 24 - 16 - 1
kcpReceiveBufSize = 128 * 1024
kcpStatsEvery = 500
kcpWindowFloor = 64
kcpWindowCeiling = 512
kcpCarrierRTT = 250 * time.Millisecond
kcpWaitSndFactor = 2
kcpBackpressurePoll = 2 * time.Millisecond
kcpChannelReliable byte = 0x00
kcpChannelRaw byte = 0x01
KCPCarrierQueueDepth = kcpWaitSndFactor * kcpWindowCeiling
)
func computeKCPWindow(fps, batch int) int {
rate := fps * batch
if rate < 1 {
rate = defaultVP8FPS * defaultVP8Batch
}
window := int(float64(rate) * kcpCarrierRTT.Seconds())
if window < kcpWindowFloor {
return kcpWindowFloor
}
if window > kcpWindowCeiling {
return kcpWindowCeiling
}
return window
}
type trackKCPSession struct {
conv uint32
vp8 *VP8DataTunnel
parent *MultiTrackKCPTunnel
kcpMu sync.Mutex
kcp *kcp.KCP
recvBuf []byte
}
func newTrackKCPSession(parent *MultiTrackKCPTunnel, vp8 *VP8DataTunnel, conv uint32, window int) *trackKCPSession {
session := &trackKCPSession{
conv: conv,
vp8: vp8,
parent: parent,
recvBuf: make([]byte, kcpReceiveBufSize),
}
session.kcp = kcp.NewKCP(conv, func(buf []byte, size int) {
if size <= 0 {
return
}
segment := make([]byte, size+1)
segment[0] = kcpChannelReliable
copy(segment[1:], buf[:size])
parent.outputSegments.Add(1)
if !session.vp8.TrySendData(segment) {
parent.droppedSegments.Add(1)
}
})
session.kcp.NoDelay(1, 10, 2, 1)
session.kcp.WndSize(window, window)
session.kcp.SetMtu(kcpSegmentMTU)
return session
}
func (s *trackKCPSession) setWindow(window int) {
s.kcpMu.Lock()
s.kcp.WndSize(window, window)
s.kcpMu.Unlock()
}
func (s *trackKCPSession) send(frame []byte) {
s.kcpMu.Lock()
s.kcp.Send(frame)
s.kcp.Update()
s.kcpMu.Unlock()
}
func (s *trackKCPSession) input(segment []byte) [][]byte {
s.kcpMu.Lock()
s.kcp.Input(segment, kcp.IKCP_PACKET_REGULAR, true)
var messages [][]byte
for {
size := s.kcp.PeekSize()
if size <= 0 {
break
}
if size > len(s.recvBuf) {
s.recvBuf = make([]byte, size)
}
n := s.kcp.Recv(s.recvBuf)
if n <= 0 {
break
}
message := make([]byte, n)
copy(message, s.recvBuf[:n])
messages = append(messages, message)
}
s.kcpMu.Unlock()
return messages
}
func (s *trackKCPSession) update() {
s.kcpMu.Lock()
s.kcp.Update()
s.kcpMu.Unlock()
}
func (s *trackKCPSession) waitSnd() int {
s.kcpMu.Lock()
pending := s.kcp.WaitSnd()
s.kcpMu.Unlock()
return pending
}
type MultiTrackKCPTunnel struct {
mt *MultiTrackTunnel
logger logger.ContextLogger
mu sync.Mutex
sessions []*trackKCPSession
convMap map[uint32]*trackKCPSession
connPin map[uint32]int
onData func([]byte)
onClose func()
stopCh chan struct{}
stopOnce sync.Once
currentWindow atomic.Int32
sentMessages atomic.Uint64
deliveredMessages atomic.Uint64
outputSegments atomic.Uint64
inputSegments atomic.Uint64
rawSent atomic.Uint64
rawReceived atomic.Uint64
droppedSegments atomic.Uint64
}
func NewMultiTrackKCPTunnel(mt *MultiTrackTunnel, logger logger.ContextLogger) *MultiTrackKCPTunnel {
t := &MultiTrackKCPTunnel{
mt: mt,
logger: logger,
convMap: make(map[uint32]*trackKCPSession),
connPin: make(map[uint32]int),
stopCh: make(chan struct{}),
}
subs := mt.SubTunnels()
window := kcpWindowFloor
if len(subs) > 0 {
window = computeKCPWindow(subs[0].FPS(), subs[0].Batch())
}
t.currentWindow.Store(int32(window))
for i, sub := range subs {
conv := uint32(kcpConvBase + i)
session := newTrackKCPSession(t, sub, conv, window)
t.sessions = append(t.sessions, session)
t.convMap[conv] = session
}
if logger != nil {
logger.Debug(fmt.Sprintf("kcptunnel: init tracks=%d window=%d queue=%d", len(subs), window, KCPCarrierQueueDepth))
}
mt.SetOnData(t.handleDecodedSegment)
mt.SetOnClose(t.handleInnerClose)
go t.updateLoop()
return t
}
func (t *MultiTrackKCPTunnel) SendData(frame []byte) {
if len(frame) < 9 {
return
}
connID := binary.BigEndian.Uint32(frame[4:8])
msgType := frame[8]
if msgType == MsgUDP || msgType == MsgUDPReply {
t.sendRaw(connID, frame)
return
}
t.mu.Lock()
if len(t.sessions) == 0 {
t.mu.Unlock()
return
}
index, pinned := t.connPin[connID]
if !pinned || index >= len(t.sessions) {
index = int(connID % uint32(len(t.sessions)))
t.connPin[connID] = index
}
session := t.sessions[index]
t.mu.Unlock()
if msgType == MsgData {
sndCap := int(t.currentWindow.Load()) * kcpWaitSndFactor
for session.waitSnd() >= sndCap {
select {
case <-t.stopCh:
return
case <-time.After(kcpBackpressurePoll):
}
}
}
t.sentMessages.Add(1)
session.send(frame)
if msgType == MsgClose {
t.mu.Lock()
delete(t.connPin, connID)
t.mu.Unlock()
}
}
func (t *MultiTrackKCPTunnel) sendRaw(connID uint32, frame []byte) {
t.mu.Lock()
if len(t.sessions) == 0 {
t.mu.Unlock()
return
}
index := int(connID % uint32(len(t.sessions)))
session := t.sessions[index]
t.mu.Unlock()
segment := make([]byte, len(frame)+1)
segment[0] = kcpChannelRaw
copy(segment[1:], frame)
t.rawSent.Add(1)
session.vp8.TrySendData(segment)
}
func (t *MultiTrackKCPTunnel) InjectSegment(payload []byte) {
t.handleDecodedSegment(payload)
}
func (t *MultiTrackKCPTunnel) handleDecodedSegment(payload []byte) {
if len(payload) < 1 {
return
}
channel := payload[0]
body := payload[1:]
if channel == kcpChannelRaw {
t.mu.Lock()
callback := t.onData
t.mu.Unlock()
if callback == nil {
return
}
t.rawReceived.Add(1)
callback(body)
return
}
if len(body) < 4 {
return
}
conv := binary.LittleEndian.Uint32(body[0:4])
t.mu.Lock()
session := t.convMap[conv]
callback := t.onData
t.mu.Unlock()
if session == nil {
return
}
t.inputSegments.Add(1)
messages := session.input(body)
if callback == nil {
return
}
for _, message := range messages {
t.deliveredMessages.Add(1)
callback(message)
}
}
func (t *MultiTrackKCPTunnel) SetOnData(fn func([]byte)) {
t.mu.Lock()
t.onData = fn
t.mu.Unlock()
}
func (t *MultiTrackKCPTunnel) SetOnClose(fn func()) {
t.mu.Lock()
t.onClose = fn
t.mu.Unlock()
}
func (t *MultiTrackKCPTunnel) Reconfigure(fps, batch int) {
t.mt.Reconfigure(fps, batch)
window := computeKCPWindow(fps, batch)
t.applyWindow(window)
if t.logger != nil {
t.logger.Debug(fmt.Sprintf("kcptunnel: reconfigure fps=%d batch=%d -> window=%d", fps, batch, window))
}
}
func (t *MultiTrackKCPTunnel) applyWindow(window int) {
t.currentWindow.Store(int32(window))
t.mu.Lock()
sessions := make([]*trackKCPSession, len(t.sessions))
copy(sessions, t.sessions)
t.mu.Unlock()
for _, session := range sessions {
session.setWindow(window)
}
}
func (t *MultiTrackKCPTunnel) AddSession(sub *VP8DataTunnel) {
window := int(t.currentWindow.Load())
t.mu.Lock()
conv := uint32(kcpConvBase + len(t.sessions))
session := newTrackKCPSession(t, sub, conv, window)
t.sessions = append(t.sessions, session)
t.convMap[conv] = session
t.mu.Unlock()
}
func (t *MultiTrackKCPTunnel) RemoveLastSession() {
t.mu.Lock()
if len(t.sessions) <= 1 {
t.mu.Unlock()
return
}
last := t.sessions[len(t.sessions)-1]
t.sessions = t.sessions[:len(t.sessions)-1]
delete(t.convMap, last.conv)
t.mu.Unlock()
}
func (t *MultiTrackKCPTunnel) Stop() {
t.stopOnce.Do(func() { close(t.stopCh) })
t.mt.Stop()
}
func (t *MultiTrackKCPTunnel) StopLayer() {
t.stopOnce.Do(func() { close(t.stopCh) })
}
func (t *MultiTrackKCPTunnel) handleInnerClose() {
t.stopOnce.Do(func() { close(t.stopCh) })
t.mu.Lock()
callback := t.onClose
t.mu.Unlock()
if callback != nil {
callback()
}
}
func (t *MultiTrackKCPTunnel) updateLoop() {
ticker := time.NewTicker(kcpUpdateInterval)
defer ticker.Stop()
ticks := 0
for {
select {
case <-t.stopCh:
return
case <-ticker.C:
t.mu.Lock()
sessions := make([]*trackKCPSession, len(t.sessions))
copy(sessions, t.sessions)
t.mu.Unlock()
for _, session := range sessions {
session.update()
}
ticks++
if ticks%kcpStatsEvery == 0 && t.logger != nil {
snmp := kcp.DefaultSnmp.Copy()
t.logger.Debug(fmt.Sprintf("kcptunnel: sessions=%d window=%d sent=%d delivered=%d out_segs=%d in_segs=%d raw_out=%d raw_in=%d dropped=%d",
len(sessions), t.currentWindow.Load(), t.sentMessages.Load(), t.deliveredMessages.Load(),
t.outputSegments.Load(), t.inputSegments.Load(),
t.rawSent.Load(), t.rawReceived.Load(), t.droppedSegments.Load()))
t.logger.Debug(fmt.Sprintf("kcptunnel: kcp_out=%d kcp_in=%d retrans=%d fastretrans=%d lost=%d repeat=%d",
snmp.OutSegs, snmp.InSegs, snmp.RetransSegs, snmp.FastRetransSegs, snmp.LostSegs, snmp.RepeatSegs))
}
}
}
}

View File

@@ -0,0 +1,199 @@
package tunnel
import (
"encoding/binary"
"sync"
)
type MultiTrackTunnel struct {
tunnels []*VP8DataTunnel
mu sync.Mutex
onData func([]byte)
onClose func()
onPeerRestart func()
isClosed bool
fps int
batch int
}
func NewMultiTrackTunnel(tunnels []*VP8DataTunnel) *MultiTrackTunnel {
m := &MultiTrackTunnel{tunnels: tunnels}
for i, tun := range tunnels {
m.wireSubTunnel(tun, i == 0)
}
return m
}
func (m *MultiTrackTunnel) AddSubTunnel(tun *VP8DataTunnel) {
m.mu.Lock()
if m.isClosed {
m.mu.Unlock()
tun.Stop()
return
}
m.tunnels = append(m.tunnels, tun)
fps := m.fps
batch := m.batch
m.mu.Unlock()
m.wireSubTunnel(tun, false)
if fps > 0 && batch > 0 {
tun.Start(fps, batch)
}
}
func (m *MultiTrackTunnel) RemoveLastSubTunnel() *VP8DataTunnel {
m.mu.Lock()
if len(m.tunnels) <= 1 {
m.mu.Unlock()
return nil
}
last := m.tunnels[len(m.tunnels)-1]
m.tunnels = m.tunnels[:len(m.tunnels)-1]
m.mu.Unlock()
last.Stop()
return last
}
func (m *MultiTrackTunnel) SubTunnelCount() int {
m.mu.Lock()
defer m.mu.Unlock()
return len(m.tunnels)
}
func (m *MultiTrackTunnel) SendData(data []byte) {
m.mu.Lock()
tunnels := m.tunnels
m.mu.Unlock()
if len(tunnels) == 0 {
return
}
var connID uint32
if len(data) >= 8 {
connID = binary.BigEndian.Uint32(data[4:8])
}
idx := connID % uint32(len(tunnels))
tunnels[idx].SendData(data)
}
func (m *MultiTrackTunnel) DeliverData(data []byte) {
m.mu.Lock()
handler := m.onData
m.mu.Unlock()
if handler != nil {
handler(data)
}
}
func (m *MultiTrackTunnel) SubTunnels() []*VP8DataTunnel {
m.mu.Lock()
defer m.mu.Unlock()
subs := make([]*VP8DataTunnel, len(m.tunnels))
copy(subs, m.tunnels)
return subs
}
func (m *MultiTrackTunnel) SetOnData(fn func([]byte)) {
m.mu.Lock()
defer m.mu.Unlock()
m.onData = fn
}
func (m *MultiTrackTunnel) SetOnClose(fn func()) {
m.mu.Lock()
defer m.mu.Unlock()
m.onClose = fn
}
func (m *MultiTrackTunnel) SetOnPeerRestart(fn func()) {
m.mu.Lock()
defer m.mu.Unlock()
m.onPeerRestart = fn
}
func (m *MultiTrackTunnel) Reconfigure(fps, batch int) {
m.mu.Lock()
m.fps = fps
m.batch = batch
tunnels := m.tunnels
m.mu.Unlock()
for _, tun := range tunnels {
tun.Reconfigure(fps, batch)
}
}
func (m *MultiTrackTunnel) Start(fps, batch int) {
m.mu.Lock()
m.fps = fps
m.batch = batch
tunnels := m.tunnels
m.mu.Unlock()
for _, tun := range tunnels {
tun.Start(fps, batch)
}
}
func (m *MultiTrackTunnel) Stop() {
m.mu.Lock()
if m.isClosed {
m.mu.Unlock()
return
}
m.isClosed = true
tunnels := m.tunnels
m.mu.Unlock()
for _, tun := range tunnels {
tun.Stop()
}
}
func (m *MultiTrackTunnel) HandleFrame(frame []byte) {
m.mu.Lock()
var first *VP8DataTunnel
if len(m.tunnels) > 0 {
first = m.tunnels[0]
}
m.mu.Unlock()
if first != nil {
first.HandleFrame(frame)
}
}
func (m *MultiTrackTunnel) wireSubTunnel(tun *VP8DataTunnel, isCamera bool) {
tun.SetOnData(func(data []byte) {
m.mu.Lock()
handler := m.onData
m.mu.Unlock()
if handler != nil {
handler(data)
}
})
if !isCamera {
return
}
tun.SetOnPeerRestart(func() {
m.mu.Lock()
handler := m.onPeerRestart
m.mu.Unlock()
if handler != nil {
handler()
}
})
tun.SetOnClose(func() {
m.mu.Lock()
if m.isClosed {
m.mu.Unlock()
return
}
m.isClosed = true
closeHandler := m.onClose
subTunnels := m.tunnels
m.mu.Unlock()
for _, t := range subTunnels {
t.Stop()
}
if closeHandler != nil {
closeHandler()
}
})
}

View File

@@ -0,0 +1,221 @@
package tunnel
import (
"crypto/cipher"
"crypto/rand"
"crypto/sha256"
"encoding/binary"
"errors"
"strings"
"sync"
"golang.org/x/crypto/chacha20poly1305"
)
var vp8Keepalive = []byte{
0x30, 0x01, 0x00, 0x9d, 0x01, 0x2a, 0x10, 0x00,
0x10, 0x00, 0x00, 0x47, 0x08, 0x85, 0x85, 0x88,
0x99, 0x84, 0x88, 0xfc,
}
var vp8Interframe = []byte{
0xb1, 0x01, 0x00, 0x08, 0x11, 0x18, 0x00, 0x18,
0x00, 0x18, 0x58, 0x2f, 0xf4, 0x00, 0x08, 0x00,
0x00,
}
const (
vp8KeepaliveLen = 20
vp8InterframeLen = 17
epochFieldLen = 4
keepaliveHdrLen = vp8KeepaliveLen + epochFieldLen
interframeHdrLen = vp8InterframeLen + epochFieldLen
)
var ErrEmptySecret = errors.New("tunnel: obfuscator requires a non-empty secret")
type DecodeResult struct {
HasFrame bool
Keepalive bool
SelfEcho bool
PeerRestart bool
Payload []byte
PeerEpoch uint32
}
type TunnelObfuscator struct {
aead cipher.AEAD
localEpoch uint32
mu sync.Mutex
peerEpoch uint32
hasPeer bool
}
func DeriveSecretFromJoinLink(joinLink string) []byte {
token := extractJoinToken(joinLink)
if token == "" {
return nil
}
return []byte(token)
}
func NewTunnelObfuscator(secret []byte) (*TunnelObfuscator, error) {
if len(secret) == 0 {
return nil, ErrEmptySecret
}
keyHash := sha256.Sum256(secret)
aead, err := chacha20poly1305.NewX(keyHash[:])
if err != nil {
return nil, err
}
var epochBytes [4]byte
if _, err := rand.Read(epochBytes[:]); err != nil {
return nil, err
}
epoch := binary.BigEndian.Uint32(epochBytes[:])
if epoch == 0 {
epoch = 1
}
return &TunnelObfuscator{aead: aead, localEpoch: epoch}, nil
}
func (o *TunnelObfuscator) LocalEpoch() uint32 { return o.localEpoch }
func (o *TunnelObfuscator) EncodeKeepalive(padLen int) []byte {
hdr := o.keepaliveHeader()
if padLen <= 0 {
return hdr
}
out := make([]byte, keepaliveHdrLen+padLen)
copy(out, hdr)
if _, err := rand.Read(out[keepaliveHdrLen:]); err != nil {
return hdr
}
return out
}
func (o *TunnelObfuscator) EncodeData(payload []byte) []byte {
hdr := o.dataHeader()
nonce := make([]byte, o.aead.NonceSize())
if _, err := rand.Read(nonce); err != nil {
return nil
}
out := make([]byte, 0, len(hdr)+len(nonce)+len(payload)+o.aead.Overhead())
out = append(out, hdr...)
out = append(out, nonce...)
out = o.aead.Seal(out, nonce, payload, nil)
return out
}
func (o *TunnelObfuscator) EncryptPayload(plaintext []byte) []byte {
if o == nil {
return plaintext
}
nonce := make([]byte, o.aead.NonceSize())
if _, err := rand.Read(nonce); err != nil {
return nil
}
out := make([]byte, 0, len(nonce)+len(plaintext)+o.aead.Overhead())
out = append(out, nonce...)
return o.aead.Seal(out, nonce, plaintext, nil)
}
func (o *TunnelObfuscator) DecryptPayload(data []byte) ([]byte, bool) {
if o == nil {
return data, true
}
nonceSize := o.aead.NonceSize()
if len(data) < nonceSize+o.aead.Overhead() {
return nil, false
}
nonce := data[:nonceSize]
ciphertext := data[nonceSize:]
plaintext, err := o.aead.Open(nil, nonce, ciphertext, nil)
if err != nil {
return nil, false
}
return plaintext, true
}
func (o *TunnelObfuscator) Decode(frame []byte) DecodeResult {
if len(frame) < 1 {
return DecodeResult{}
}
var hdrLen, epochOff int
isKeepaliveFrame := false
switch frame[0] {
case vp8Keepalive[0]:
hdrLen = keepaliveHdrLen
epochOff = vp8KeepaliveLen
isKeepaliveFrame = true
case vp8Interframe[0]:
hdrLen = interframeHdrLen
epochOff = vp8InterframeLen
default:
return DecodeResult{}
}
if len(frame) < hdrLen {
return DecodeResult{}
}
peerEpoch := binary.BigEndian.Uint32(frame[epochOff : epochOff+epochFieldLen])
if peerEpoch == o.localEpoch {
return DecodeResult{HasFrame: true, SelfEcho: true, PeerEpoch: peerEpoch}
}
res := DecodeResult{HasFrame: true, PeerEpoch: peerEpoch}
o.mu.Lock()
if !o.hasPeer {
o.peerEpoch = peerEpoch
o.hasPeer = true
} else if o.peerEpoch != peerEpoch {
o.peerEpoch = peerEpoch
res.PeerRestart = true
}
o.mu.Unlock()
if isKeepaliveFrame || len(frame) == hdrLen {
res.Keepalive = true
return res
}
body := frame[hdrLen:]
nonceSize := o.aead.NonceSize()
if len(body) < nonceSize+o.aead.Overhead() {
return DecodeResult{}
}
nonce := body[:nonceSize]
ciphertext := body[nonceSize:]
plaintext, err := o.aead.Open(nil, nonce, ciphertext, nil)
if err != nil {
return DecodeResult{}
}
res.Payload = plaintext
return res
}
func extractJoinToken(joinLink string) string {
s := strings.TrimSpace(joinLink)
s = strings.TrimRight(s, "/")
if i := strings.IndexByte(s, '?'); i >= 0 {
s = s[:i]
}
if i := strings.IndexByte(s, '#'); i >= 0 {
s = s[:i]
}
if i := strings.LastIndexByte(s, '/'); i >= 0 {
s = s[i+1:]
}
return s
}
func (o *TunnelObfuscator) keepaliveHeader() []byte {
hdr := make([]byte, keepaliveHdrLen)
copy(hdr, vp8Keepalive)
binary.BigEndian.PutUint32(hdr[vp8KeepaliveLen:], o.localEpoch)
return hdr
}
func (o *TunnelObfuscator) dataHeader() []byte {
hdr := make([]byte, interframeHdrLen)
copy(hdr, vp8Interframe)
binary.BigEndian.PutUint32(hdr[vp8InterframeLen:], o.localEpoch)
return hdr
}

View File

@@ -0,0 +1,94 @@
package tunnel
import "encoding/binary"
const (
MsgConnect byte = 0x01
MsgConnectOK byte = 0x02
MsgConnectErr byte = 0x03
MsgData byte = 0x04
MsgClose byte = 0x05
MsgUDP byte = 0x06
MsgUDPReply byte = 0x07
MsgConfig byte = 0x08
MsgConfigAck byte = 0x09
)
const ControlConnID uint32 = 0
type DataTunnel interface {
SendData(data []byte)
SetOnData(fn func([]byte))
SetOnClose(fn func())
Reconfigure(fps, batch int)
}
func EncodeVP8Config(fps, batch, trackCount int) []byte {
if fps < 1 {
fps = 1
}
if batch < 1 {
batch = 1
}
if trackCount < 1 {
trackCount = 1
}
if fps > 0xFFFF {
fps = 0xFFFF
}
if batch > 0xFFFF {
batch = 0xFFFF
}
if trackCount > 0xFFFF {
trackCount = 0xFFFF
}
var payload [6]byte
binary.BigEndian.PutUint16(payload[0:2], uint16(fps))
binary.BigEndian.PutUint16(payload[2:4], uint16(batch))
binary.BigEndian.PutUint16(payload[4:6], uint16(trackCount))
return EncodeFrame(ControlConnID, MsgConfig, payload[:])
}
func DecodeVP8Config(payload []byte) (fps, batch, trackCount int, ok bool) {
if len(payload) < 4 {
return 0, 0, 0, false
}
fps = int(binary.BigEndian.Uint16(payload[0:2]))
batch = int(binary.BigEndian.Uint16(payload[2:4]))
trackCount = 1
if len(payload) >= 6 {
trackCount = int(binary.BigEndian.Uint16(payload[4:6]))
}
return fps, batch, trackCount, true
}
func EncodeFrame(connID uint32, msgType byte, payload []byte) []byte {
buf := make([]byte, 4+5+len(payload))
binary.BigEndian.PutUint32(buf[0:4], uint32(5+len(payload)))
binary.BigEndian.PutUint32(buf[4:8], connID)
buf[8] = msgType
copy(buf[9:], payload)
return buf
}
func LooksLikeRelayFrame(payload []byte) bool {
if len(payload) < 9 {
return false
}
frameLen := binary.BigEndian.Uint32(payload[0:4])
return frameLen >= 5 && int(frameLen)+4 <= len(payload)
}
func DecodeFrames(data []byte, cb func(connID uint32, msgType byte, payload []byte)) {
for len(data) >= 4 {
frameLen := int(binary.BigEndian.Uint32(data[0:4]))
if frameLen < 5 || 4+frameLen > len(data) {
return
}
connID := binary.BigEndian.Uint32(data[4:8])
msgType := data[8]
payload := data[9 : 4+frameLen]
cb(connID, msgType, payload)
data = data[4+frameLen:]
}
}

View File

@@ -0,0 +1,662 @@
package tunnel
import (
"bytes"
"context"
"fmt"
"io"
"net"
"sync"
"sync/atomic"
"time"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
type udpClient struct {
pending chan []byte
closed atomic.Bool
addr string
}
type RelayBridge struct {
tunnelMu sync.RWMutex
tunnel DataTunnel
conns sync.Map
udpClients sync.Map
nextID atomic.Uint32
logger logger.ContextLogger
mode string
readBuf int
ready chan struct{}
once sync.Once
closed atomic.Bool
dialer N.Dialer
acceptHandlerMu sync.Mutex
acceptHandler func(conn net.Conn, destination string)
udpAcceptHandlerMu sync.Mutex
udpAcceptHandler func(conn net.Conn, destination string)
onPeerConfigMu sync.Mutex
onPeerConfig func(fps, batch, trackCount int)
}
func NewRelayBridge(tunnel DataTunnel, mode string, readBuf int, dialer N.Dialer, logger logger.ContextLogger) *RelayBridge {
rb := &RelayBridge{
tunnel: tunnel,
logger: logger,
mode: mode,
readBuf: readBuf,
dialer: dialer,
ready: make(chan struct{}),
}
tunnel.SetOnData(rb.handleTunnelData)
tunnel.SetOnClose(rb.handleTunnelClose)
return rb
}
func (rb *RelayBridge) SetAcceptHandler(fn func(conn net.Conn, destination string)) {
rb.acceptHandlerMu.Lock()
rb.acceptHandler = fn
rb.acceptHandlerMu.Unlock()
}
func (rb *RelayBridge) SetUDPAcceptHandler(fn func(conn net.Conn, destination string)) {
rb.udpAcceptHandlerMu.Lock()
rb.udpAcceptHandler = fn
rb.udpAcceptHandlerMu.Unlock()
}
func (rb *RelayBridge) SetOnPeerConfig(fn func(fps, batch, trackCount int)) {
rb.onPeerConfigMu.Lock()
rb.onPeerConfig = fn
rb.onPeerConfigMu.Unlock()
}
func (rb *RelayBridge) DialContext(ctx context.Context, destination string) (net.Conn, error) {
if rb.closed.Load() {
return nil, fmt.Errorf("relay: bridge already closed")
}
if M.ParseSocksaddr(destination).IsIPv6() {
return nil, fmt.Errorf("relay: network unreachable (ipv6): %s", common.MaskAddr(destination))
}
select {
case <-rb.ready:
case <-ctx.Done():
return nil, ctx.Err()
}
id := rb.nextID.Add(1)
tc := newTunnelConn(id, rb)
rb.conns.Store(id, tc)
rb.logger.Debug(fmt.Sprintf("relay: DIAL %d -> %s", id, common.MaskAddr(destination)))
rb.send(id, MsgConnect, []byte(destination))
select {
case err := <-tc.rdy:
if err != nil {
rb.conns.Delete(id)
return nil, err
}
return tc, nil
case <-ctx.Done():
rb.conns.Delete(id)
rb.send(id, MsgClose, nil)
return nil, ctx.Err()
}
}
func (rb *RelayBridge) ListenPacket(ctx context.Context, destination string) (net.Conn, error) {
if rb.closed.Load() {
return nil, fmt.Errorf("relay: bridge already closed")
}
if M.ParseSocksaddr(destination).IsIPv6() {
return nil, fmt.Errorf("relay: network unreachable (ipv6): %s", common.MaskAddr(destination))
}
select {
case <-rb.ready:
case <-ctx.Done():
return nil, ctx.Err()
}
id := rb.nextID.Add(1)
uc := &udpClient{pending: make(chan []byte, 64), addr: destination}
rb.udpClients.Store(id, uc)
return &tunnelPacketConn{id: id, rb: rb, uc: uc, destStr: destination}, nil
}
func (rb *RelayBridge) Reset() {
rb.closeAll()
}
func (rb *RelayBridge) Close() {
if !rb.closed.CompareAndSwap(false, true) {
return
}
rb.closeAll()
}
func (rb *RelayBridge) MarkReady() {
rb.once.Do(func() { close(rb.ready) })
}
func (rb *RelayBridge) currentTunnel() DataTunnel {
rb.tunnelMu.RLock()
defer rb.tunnelMu.RUnlock()
return rb.tunnel
}
func (rb *RelayBridge) SwapTunnel(newTunnel DataTunnel) {
rb.tunnelMu.Lock()
rb.tunnel = newTunnel
rb.tunnelMu.Unlock()
newTunnel.SetOnData(rb.handleTunnelData)
newTunnel.SetOnClose(rb.handleTunnelClose)
rb.closeAll()
}
func (rb *RelayBridge) IsClosed() bool {
return rb.closed.Load()
}
func (rb *RelayBridge) handleTunnelClose() {
rb.closeAll()
}
func (rb *RelayBridge) closeAll() {
var ids []uint32
rb.conns.Range(func(key, value any) bool {
if id, ok := key.(uint32); ok {
ids = append(ids, id)
}
if c, ok := value.(net.Conn); ok {
c.Close()
}
rb.conns.Delete(key)
return true
})
udpCount := 0
rb.udpClients.Range(func(key, value any) bool {
udpCount++
if uc, ok := value.(*udpClient); ok {
uc.closed.Store(true)
close(uc.pending)
}
rb.udpClients.Delete(key)
return true
})
rb.logger.Debug(fmt.Sprintf("relay: closeAll mode=%s tcp=%d udp=%d ids=%v nextID=%d", rb.mode, len(ids), udpCount, ids, rb.nextID.Load()))
}
func (rb *RelayBridge) send(connID uint32, msgType byte, payload []byte) {
frame := EncodeFrame(connID, msgType, payload)
rb.currentTunnel().SendData(frame)
}
func (rb *RelayBridge) handleTunnelData(data []byte) {
DecodeFrames(data, func(connID uint32, msgType byte, payload []byte) {
if connID == ControlConnID && msgType == MsgConfig {
fps, batch, trackCount, ok := DecodeVP8Config(payload)
if !ok {
return
}
if rb.mode == "creator" {
rb.logger.Debug(fmt.Sprintf("relay: peer requested vp8 pacing fps=%d batch=%d trackCount=%d", fps, batch, trackCount))
rb.currentTunnel().Reconfigure(fps, batch)
rb.send(ControlConnID, MsgConfigAck, nil)
rb.onPeerConfigMu.Lock()
cb := rb.onPeerConfig
rb.onPeerConfigMu.Unlock()
if cb != nil {
cb(fps, batch, trackCount)
}
}
return
}
if connID == ControlConnID && msgType == MsgConfigAck {
return
}
switch rb.mode {
case "joiner":
rb.handleJoinerMessage(connID, msgType, payload)
case "creator":
rb.handleCreatorMessage(connID, msgType, payload)
}
})
}
func (rb *RelayBridge) handleJoinerMessage(connID uint32, msgType byte, payload []byte) {
if msgType == MsgUDPReply {
uval, ok := rb.udpClients.Load(connID)
if !ok {
return
}
uc := uval.(*udpClient)
if uc.closed.Load() {
return
}
cp := make([]byte, len(payload))
copy(cp, payload)
select {
case uc.pending <- cp:
default:
}
return
}
val, ok := rb.conns.Load(connID)
if !ok {
if msgType != MsgClose {
rb.logger.Debug(fmt.Sprintf("relay[joiner]: drop msgType=%d for unknown conn %d (payload=%dB)", msgType, connID, len(payload)))
}
return
}
tc := val.(*tunnelConn)
switch msgType {
case MsgConnectOK:
select {
case tc.rdy <- nil:
default:
}
case MsgConnectErr:
select {
case tc.rdy <- fmt.Errorf("%s", payload):
default:
}
case MsgData:
tc.deliver(payload)
case MsgClose:
tc.remoteClosed()
rb.conns.Delete(connID)
}
}
func (rb *RelayBridge) handleCreatorMessage(connID uint32, msgType byte, payload []byte) {
switch msgType {
case MsgConnect:
rb.acceptHandlerMu.Lock()
handler := rb.acceptHandler
rb.acceptHandlerMu.Unlock()
if handler != nil {
destination := string(payload)
tc := newTunnelConn(connID, rb)
rb.conns.Store(connID, tc)
rb.send(connID, MsgConnectOK, nil)
go handler(tc, destination)
return
}
go rb.connectTCP(connID, string(payload))
case MsgUDP:
payloadCopy := make([]byte, len(payload))
copy(payloadCopy, payload)
go rb.handleUDP(connID, payloadCopy)
case MsgData:
val, ok := rb.conns.Load(connID)
if !ok {
rb.logger.Debug(fmt.Sprintf("relay[creator]: drop MsgData for unknown conn %d (payload=%dB)", connID, len(payload)))
rb.send(connID, MsgClose, nil)
return
}
switch c := val.(type) {
case *tunnelConn:
c.deliver(payload)
case net.Conn:
if _, err := c.Write(payload); err != nil {
rb.logger.Debug(fmt.Sprintf("relay[creator]: write to target %d failed: %s", connID, common.MaskError(err)))
}
}
case MsgClose:
found := false
if val, ok := rb.conns.LoadAndDelete(connID); ok {
found = true
switch c := val.(type) {
case *tunnelConn:
c.remoteClosed()
case net.Conn:
c.Close()
}
}
if uval, ok := rb.udpClients.LoadAndDelete(connID); ok {
found = true
switch uc := uval.(type) {
case *creatorUDPConn:
uc.remoteClosed()
case net.Conn:
uc.Close()
}
}
if !found {
rb.logger.Debug(fmt.Sprintf("relay[creator]: drop MsgClose for unknown conn %d", connID))
}
}
}
func (rb *RelayBridge) handleUDP(connID uint32, payload []byte) {
if len(payload) < 2 {
return
}
addrLen := int(payload[0])
if addrLen == 0 || len(payload) < 1+addrLen {
return
}
if bytes.IndexByte(payload[1:1+addrLen], 0) != -1 {
return
}
addr := string(payload[1 : 1+addrLen])
data := payload[1+addrLen:]
rb.udpAcceptHandlerMu.Lock()
handler := rb.udpAcceptHandler
rb.udpAcceptHandlerMu.Unlock()
if handler != nil {
var cuc *creatorUDPConn
if val, ok := rb.udpClients.Load(connID); ok {
existing, ok := val.(*creatorUDPConn)
if !ok {
return
}
cuc = existing
} else {
created := newCreatorUDPConn(connID, rb, addr)
if actual, loaded := rb.udpClients.LoadOrStore(connID, created); loaded {
existing, ok := actual.(*creatorUDPConn)
if !ok {
return
}
cuc = existing
} else {
cuc = created
go handler(cuc, addr)
}
}
cuc.deliver(data)
return
}
var egress net.Conn
if val, ok := rb.udpClients.Load(connID); ok {
existing, ok := val.(net.Conn)
if !ok {
return
}
egress = existing
} else {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
created, err := rb.dialer.DialContext(ctx, N.NetworkUDP, M.ParseSocksaddr(addr))
cancel()
if err != nil {
rb.logger.Warn(fmt.Sprintf("relay[creator]: UDP %d open %s failed: %v", connID, common.MaskAddr(addr), err))
return
}
if actual, loaded := rb.udpClients.LoadOrStore(connID, created); loaded {
created.Close()
existing, ok := actual.(net.Conn)
if !ok {
return
}
egress = existing
} else {
egress = created
go func(conn net.Conn, id uint32, target string) {
defer conn.Close()
defer rb.udpClients.Delete(id)
defer rb.send(id, MsgClose, nil)
buf := make([]byte, common.UDPBufSize)
for {
conn.SetReadDeadline(time.Now().Add(60 * time.Second))
n, err := conn.Read(buf)
if err != nil {
return
}
rb.send(id, MsgUDPReply, buf[:n])
}
}(egress, connID, addr)
}
}
egress.SetWriteDeadline(time.Now().Add(5 * time.Second))
if _, err := egress.Write(data); err != nil {
rb.logger.Debug(fmt.Sprintf("relay[creator]: UDP %d write %s failed: %v", connID, common.MaskAddr(addr), err))
}
}
func (rb *RelayBridge) connectTCP(connID uint32, addr string) {
rb.logger.Debug(fmt.Sprintf("relay: CONNECT %d -> %s", connID, common.MaskAddr(addr)))
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
conn, err := rb.dialer.DialContext(ctx, N.NetworkTCP, M.ParseSocksaddr(addr))
cancel()
if err != nil {
rb.logger.Warn(fmt.Sprintf("relay: CONNECT %d failed: %s", connID, common.MaskError(err)))
rb.send(connID, MsgConnectErr, []byte(common.MaskError(err)))
return
}
rb.conns.Store(connID, conn)
rb.send(connID, MsgConnectOK, nil)
rb.logger.Debug(fmt.Sprintf("relay: CONNECTED %d -> %s", connID, common.MaskAddr(addr)))
buf := make([]byte, rb.readBuf)
var totalRead int64
var reads int
for {
n, err := conn.Read(buf)
if n > 0 {
rb.send(connID, MsgData, buf[:n])
totalRead += int64(n)
reads++
}
if err != nil {
if err != io.EOF {
rb.logger.Warn(fmt.Sprintf("relay: conn %d read error: %s (read %d times, %dB)", connID, common.MaskError(err), reads, totalRead))
}
break
}
}
rb.send(connID, MsgClose, nil)
rb.conns.Delete(connID)
}
type tunnelAddr struct{}
func (tunnelAddr) Network() string { return "call" }
func (tunnelAddr) String() string { return "call" }
type tunnelConn struct {
id uint32
rb *RelayBridge
rdy chan error
readBuf bytes.Buffer
readMu sync.Mutex
readCond chan struct{}
closed atomic.Bool
closeCh chan struct{}
}
func newTunnelConn(id uint32, rb *RelayBridge) *tunnelConn {
return &tunnelConn{
id: id,
rb: rb,
rdy: make(chan error, 1),
readCond: make(chan struct{}, 1),
closeCh: make(chan struct{}),
}
}
func (tc *tunnelConn) Read(b []byte) (int, error) {
for {
tc.readMu.Lock()
if tc.readBuf.Len() > 0 {
n, _ := tc.readBuf.Read(b)
tc.readMu.Unlock()
return n, nil
}
tc.readMu.Unlock()
select {
case <-tc.closeCh:
tc.readMu.Lock()
if tc.readBuf.Len() > 0 {
n, _ := tc.readBuf.Read(b)
tc.readMu.Unlock()
return n, nil
}
tc.readMu.Unlock()
return 0, io.EOF
case <-tc.readCond:
}
}
}
func (tc *tunnelConn) Write(b []byte) (int, error) {
if tc.closed.Load() {
return 0, io.ErrClosedPipe
}
tc.rb.send(tc.id, MsgData, b)
return len(b), nil
}
func (tc *tunnelConn) Close() error {
if tc.closed.CompareAndSwap(false, true) {
close(tc.closeCh)
tc.rb.send(tc.id, MsgClose, nil)
tc.rb.conns.Delete(tc.id)
}
return nil
}
func (tc *tunnelConn) LocalAddr() net.Addr { return tunnelAddr{} }
func (tc *tunnelConn) RemoteAddr() net.Addr { return tunnelAddr{} }
func (tc *tunnelConn) SetDeadline(t time.Time) error { return nil }
func (tc *tunnelConn) SetReadDeadline(t time.Time) error { return nil }
func (tc *tunnelConn) SetWriteDeadline(t time.Time) error { return nil }
func (tc *tunnelConn) deliver(payload []byte) {
tc.readMu.Lock()
tc.readBuf.Write(payload)
tc.readMu.Unlock()
select {
case tc.readCond <- struct{}{}:
default:
}
}
func (tc *tunnelConn) remoteClosed() {
if tc.closed.CompareAndSwap(false, true) {
close(tc.closeCh)
}
}
type tunnelPacketConn struct {
id uint32
rb *RelayBridge
uc *udpClient
destStr string
}
func (pc *tunnelPacketConn) Read(b []byte) (int, error) {
data, ok := <-pc.uc.pending
if !ok {
return 0, io.EOF
}
n := copy(b, data)
return n, nil
}
func (pc *tunnelPacketConn) Write(b []byte) (int, error) {
if pc.uc.closed.Load() {
return 0, io.ErrClosedPipe
}
payload := make([]byte, 1+len(pc.destStr)+len(b))
payload[0] = byte(len(pc.destStr))
copy(payload[1:], pc.destStr)
copy(payload[1+len(pc.destStr):], b)
pc.rb.send(pc.id, MsgUDP, payload)
return len(b), nil
}
func (pc *tunnelPacketConn) Close() error {
if pc.uc.closed.CompareAndSwap(false, true) {
close(pc.uc.pending)
pc.rb.udpClients.Delete(pc.id)
pc.rb.send(pc.id, MsgClose, nil)
}
return nil
}
func (pc *tunnelPacketConn) LocalAddr() net.Addr { return tunnelAddr{} }
func (pc *tunnelPacketConn) RemoteAddr() net.Addr { return tunnelAddr{} }
func (pc *tunnelPacketConn) SetDeadline(t time.Time) error { return nil }
func (pc *tunnelPacketConn) SetReadDeadline(t time.Time) error { return nil }
func (pc *tunnelPacketConn) SetWriteDeadline(t time.Time) error { return nil }
type creatorUDPConn struct {
id uint32
rb *RelayBridge
addr string
readBuf bytes.Buffer
readMu sync.Mutex
readCond chan struct{}
closed atomic.Bool
closeCh chan struct{}
}
func newCreatorUDPConn(id uint32, rb *RelayBridge, addr string) *creatorUDPConn {
return &creatorUDPConn{
id: id,
rb: rb,
addr: addr,
readCond: make(chan struct{}, 1),
closeCh: make(chan struct{}),
}
}
func (uc *creatorUDPConn) Read(b []byte) (int, error) {
for {
uc.readMu.Lock()
if uc.readBuf.Len() > 0 {
n, _ := uc.readBuf.Read(b)
uc.readMu.Unlock()
return n, nil
}
uc.readMu.Unlock()
select {
case <-uc.closeCh:
return 0, io.EOF
case <-uc.readCond:
}
}
}
func (uc *creatorUDPConn) Write(b []byte) (int, error) {
if uc.closed.Load() {
return 0, io.ErrClosedPipe
}
uc.rb.send(uc.id, MsgUDPReply, b)
return len(b), nil
}
func (uc *creatorUDPConn) Close() error {
if uc.closed.CompareAndSwap(false, true) {
close(uc.closeCh)
uc.rb.send(uc.id, MsgClose, nil)
uc.rb.udpClients.Delete(uc.id)
}
return nil
}
func (uc *creatorUDPConn) LocalAddr() net.Addr { return tunnelAddr{} }
func (uc *creatorUDPConn) RemoteAddr() net.Addr { return tunnelAddr{} }
func (uc *creatorUDPConn) SetDeadline(t time.Time) error { return nil }
func (uc *creatorUDPConn) SetReadDeadline(t time.Time) error { return nil }
func (uc *creatorUDPConn) SetWriteDeadline(t time.Time) error { return nil }
func (uc *creatorUDPConn) deliver(payload []byte) {
uc.readMu.Lock()
uc.readBuf.Write(payload)
uc.readMu.Unlock()
select {
case uc.readCond <- struct{}{}:
default:
}
}
func (uc *creatorUDPConn) remoteClosed() {
if uc.closed.CompareAndSwap(false, true) {
close(uc.closeCh)
}
}

View File

@@ -0,0 +1,278 @@
package tunnel
import (
"encoding/binary"
"fmt"
"sync"
"sync/atomic"
"time"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing/common/logger"
)
const (
screenWriterFPS = 24
screenWriterBatch = 30
screenWriterMaxBytes = 60000
screenWriterQueue = 256
screenKeepalivePadMax = 48
)
type ScreenWriter struct {
obf *TunnelObfuscator
logger logger.ContextLogger
label string
sendMu sync.Mutex
send func([]byte) error
stopCh chan struct{}
sendQueue chan []byte
cfgChan chan struct{}
stopOnce sync.Once
running atomic.Bool
cfgMu sync.Mutex
fps int
batch int
sent atomic.Uint64
}
func NewScreenWriter(obf *TunnelObfuscator, label string, logger logger.ContextLogger) *ScreenWriter {
return &ScreenWriter{
obf: obf,
logger: logger,
label: label,
stopCh: make(chan struct{}),
sendQueue: make(chan []byte, screenWriterQueue),
cfgChan: make(chan struct{}, 1),
fps: screenWriterFPS,
batch: screenWriterBatch,
}
}
func (w *ScreenWriter) SetSend(fn func([]byte) error) {
w.sendMu.Lock()
w.send = fn
w.sendMu.Unlock()
}
func (w *ScreenWriter) SendData(data []byte) {
if len(data) == 0 {
return
}
select {
case w.sendQueue <- data:
case <-w.stopCh:
}
}
func (w *ScreenWriter) Reconfigure(fps, batch int) {
if fps <= 0 && batch <= 0 {
return
}
w.cfgMu.Lock()
changed := false
if fps > 0 && w.fps != fps {
w.fps = fps
changed = true
}
if batch > 0 && w.batch != batch {
w.batch = batch
changed = true
}
w.cfgMu.Unlock()
if changed {
select {
case w.cfgChan <- struct{}{}:
default:
}
}
}
func (w *ScreenWriter) Start() {
if !w.running.CompareAndSwap(false, true) {
return
}
go w.writerLoop()
}
func (w *ScreenWriter) Stop() {
if !w.running.CompareAndSwap(true, false) {
return
}
w.stopOnce.Do(func() { close(w.stopCh) })
}
func (w *ScreenWriter) interval() time.Duration {
w.cfgMu.Lock()
fps, batch := w.fps, w.batch
w.cfgMu.Unlock()
frame := time.Second / time.Duration(fps)
sample := frame
if batch > 1 {
sample = frame / time.Duration(batch)
}
if sample <= 0 {
sample = time.Millisecond
}
return sample
}
func (w *ScreenWriter) nextKeepalive(sample time.Duration) (ticks, padLen int) {
ticks = int(common.DurationInRange(keepaliveIdleMin, keepaliveIdleMax) / sample)
if ticks < 1 {
ticks = 1
}
return ticks, common.IntInRange(0, screenKeepalivePadMax)
}
func (w *ScreenWriter) emit(msg []byte) {
if msg == nil || len(msg) > screenWriterMaxBytes {
return
}
w.sendMu.Lock()
send := w.send
w.sendMu.Unlock()
if send == nil {
return
}
if err := send(msg); err != nil {
return
}
n := w.sent.Add(1)
if n <= 5 || n%500 == 0 {
w.logger.Debug(fmt.Sprintf("[%s] sent frame #%d size=%d", w.label, n, len(msg)))
}
}
func (w *ScreenWriter) writerLoop() {
for {
sample := w.interval()
keepaliveEvery, keepalivePad := w.nextKeepalive(sample)
ticker := time.NewTicker(sample)
idle := 0
reconfigure := false
for !reconfigure {
select {
case <-w.stopCh:
ticker.Stop()
return
case <-w.cfgChan:
reconfigure = true
case <-ticker.C:
select {
case data := <-w.sendQueue:
w.emit(w.obf.EncodeData(data))
idle = 0
default:
idle++
if idle < keepaliveEvery {
continue
}
idle = 0
w.emit(w.obf.EncodeKeepalive(keepalivePad))
keepaliveEvery, keepalivePad = w.nextKeepalive(sample)
}
}
}
ticker.Stop()
}
}
type SymmetricScreenTunnel struct {
cam *VP8DataTunnel
screen *ScreenWriter
obf *TunnelObfuscator
logger logger.ContextLogger
screenReady func() bool
onDataMu sync.Mutex
onData func([]byte)
recv atomic.Uint64
trackCount atomic.Int32
}
func NewSymmetricScreenTunnel(cam *VP8DataTunnel, screen *ScreenWriter, obf *TunnelObfuscator, screenReady func() bool, logger logger.ContextLogger) *SymmetricScreenTunnel {
return &SymmetricScreenTunnel{cam: cam, screen: screen, obf: obf, screenReady: screenReady, logger: logger}
}
func (s *SymmetricScreenTunnel) SetTrackCount(n int) {
if n < 1 {
n = 1
}
if n > 2 {
n = 2
}
old := s.trackCount.Swap(int32(n))
if int(old) != n {
s.logger.Debug(fmt.Sprintf("screen tunnel track count %d -> %d", old, n))
}
if n >= 2 {
s.screen.Start()
}
}
func (s *SymmetricScreenTunnel) SendData(data []byte) {
var connID uint32
if len(data) >= 8 {
connID = binary.BigEndian.Uint32(data[4:8])
}
if connID == ControlConnID {
s.cam.SendData(data)
return
}
tc := uint32(s.trackCount.Load())
if tc < 1 {
tc = 1
}
if connID%tc == 1 && s.screenUp() {
s.screen.SendData(data)
return
}
s.cam.SendData(data)
}
func (s *SymmetricScreenTunnel) SetOnData(fn func([]byte)) {
s.onDataMu.Lock()
s.onData = fn
s.onDataMu.Unlock()
s.cam.SetOnData(fn)
}
func (s *SymmetricScreenTunnel) SetOnClose(fn func()) { s.cam.SetOnClose(fn) }
func (s *SymmetricScreenTunnel) Reconfigure(fps, batch int) {
s.cam.Reconfigure(fps, batch)
s.screen.Reconfigure(fps, batch)
}
func (s *SymmetricScreenTunnel) Stop() {
s.screen.Stop()
s.cam.Stop()
}
func (s *SymmetricScreenTunnel) HandleScreenFrame(frame []byte) {
res := s.obf.Decode(frame)
n := s.recv.Add(1)
if n <= 10 || n%500 == 0 {
s.logger.Debug(fmt.Sprintf("screen recv frame #%d in=%d hasFrame=%v keepalive=%v payload=%d", n, len(frame), res.HasFrame, res.Keepalive, len(res.Payload)))
}
if !res.HasFrame || res.SelfEcho || res.Keepalive || len(res.Payload) == 0 {
return
}
s.onDataMu.Lock()
handler := s.onData
s.onDataMu.Unlock()
if handler != nil {
handler(res.Payload)
}
}
func (s *SymmetricScreenTunnel) screenUp() bool {
if s.screenReady == nil {
return true
}
return s.screenReady()
}

View File

@@ -0,0 +1,319 @@
package tunnel
import (
"fmt"
"sync"
"sync/atomic"
"time"
"github.com/pion/webrtc/v4"
"github.com/pion/webrtc/v4/pkg/media"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing/common/logger"
)
const (
defaultVP8FPS = 24
defaultVP8Batch = 30
keepaliveIdleMin = 60 * time.Millisecond
keepaliveIdleMax = 200 * time.Millisecond
keepalivePadMax = 176
sendQueueDepth = 128
paceBatchFloorPercent = 80
paceDriftMin = 5 * time.Second
paceDriftMax = 20 * time.Second
)
type VP8DataTunnel struct {
track *webrtc.TrackLocalStaticSample
logger logger.ContextLogger
obf *TunnelObfuscator
stopCh chan struct{}
sendQueue chan []byte
cfgChan chan struct{}
stopOnce sync.Once
running atomic.Bool
cfgMu sync.Mutex
fps int
batch int
keepaliveMin time.Duration
keepaliveMax time.Duration
keepalivePadMax int
sentFrames atomic.Uint64
recvFrames atomic.Uint64
keepaliveFrames atomic.Uint64
OnData func([]byte)
OnClose func()
OnPeerRestart func()
}
func (t *VP8DataTunnel) SetOnData(fn func([]byte)) { t.OnData = fn }
func (t *VP8DataTunnel) SetOnClose(fn func()) { t.OnClose = fn }
func (t *VP8DataTunnel) SetOnPeerRestart(fn func()) { t.OnPeerRestart = fn }
func NewVP8DataTunnel(track *webrtc.TrackLocalStaticSample, obf *TunnelObfuscator, logger logger.ContextLogger) *VP8DataTunnel {
return NewVP8DataTunnelWithQueue(track, obf, logger, sendQueueDepth)
}
func NewVP8DataTunnelWithQueue(track *webrtc.TrackLocalStaticSample, obf *TunnelObfuscator, logger logger.ContextLogger, queueDepth int) *VP8DataTunnel {
if queueDepth < sendQueueDepth {
queueDepth = sendQueueDepth
}
return &VP8DataTunnel{
track: track,
obf: obf,
logger: logger,
stopCh: make(chan struct{}),
sendQueue: make(chan []byte, queueDepth),
cfgChan: make(chan struct{}, 1),
fps: defaultVP8FPS,
batch: defaultVP8Batch,
keepaliveMin: keepaliveIdleMin,
keepaliveMax: keepaliveIdleMax,
keepalivePadMax: keepalivePadMax,
}
}
func (t *VP8DataTunnel) SetKeepaliveShape(minPeriod, maxPeriod time.Duration, padMax int) {
t.cfgMu.Lock()
if minPeriod > 0 {
t.keepaliveMin = minPeriod
}
if maxPeriod >= t.keepaliveMin {
t.keepaliveMax = maxPeriod
}
if padMax >= 0 {
t.keepalivePadMax = padMax
}
newMin, newMax, newPad := t.keepaliveMin, t.keepaliveMax, t.keepalivePadMax
t.cfgMu.Unlock()
t.logger.Debug(fmt.Sprintf("vp8tunnel: keepalive shape min=%s max=%s padMax=%d", newMin, newMax, newPad))
}
func (t *VP8DataTunnel) nextKeepalive(sampleInterval time.Duration) (ticks, padLen int) {
t.cfgMu.Lock()
minPeriod, maxPeriod, padMax := t.keepaliveMin, t.keepaliveMax, t.keepalivePadMax
t.cfgMu.Unlock()
ticks = int(common.DurationInRange(minPeriod, maxPeriod) / sampleInterval)
if ticks < 1 {
ticks = 1
}
return ticks, common.IntInRange(0, padMax)
}
func (t *VP8DataTunnel) Reconfigure(fps, batch int) {
if fps <= 0 && batch <= 0 {
return
}
t.cfgMu.Lock()
changed := false
if fps > 0 && t.fps != fps {
t.fps = fps
changed = true
}
if batch > 0 && t.batch != batch {
t.batch = batch
changed = true
}
newFPS, newBatch := t.fps, t.batch
t.cfgMu.Unlock()
if !changed {
return
}
t.logger.Debug(fmt.Sprintf("vp8tunnel: reconfigure fps=%d batch=%d", newFPS, newBatch))
select {
case t.cfgChan <- struct{}{}:
default:
}
}
func (t *VP8DataTunnel) FPS() int {
t.cfgMu.Lock()
defer t.cfgMu.Unlock()
return t.fps
}
func (t *VP8DataTunnel) Batch() int {
t.cfgMu.Lock()
defer t.cfgMu.Unlock()
return t.batch
}
func (t *VP8DataTunnel) SendData(data []byte) {
if len(data) == 0 {
return
}
select {
case t.sendQueue <- data:
case <-t.stopCh:
}
}
func (t *VP8DataTunnel) TrySendData(data []byte) bool {
if len(data) == 0 {
return true
}
select {
case t.sendQueue <- data:
return true
case <-t.stopCh:
return false
default:
return false
}
}
func (t *VP8DataTunnel) Start(fps, batch int) {
t.cfgMu.Lock()
if fps > 0 {
t.fps = fps
}
if batch > 0 {
t.batch = batch
}
t.cfgMu.Unlock()
if !t.running.CompareAndSwap(false, true) {
return
}
go t.writerLoop()
}
func (t *VP8DataTunnel) Stop() {
if !t.running.CompareAndSwap(true, false) {
return
}
t.stopOnce.Do(func() { close(t.stopCh) })
if t.OnClose != nil {
t.OnClose()
}
}
func (t *VP8DataTunnel) HandleFrame(frame []byte) {
res := t.obf.Decode(frame)
if !res.HasFrame {
return
}
if res.SelfEcho {
return
}
if res.PeerRestart {
t.logger.Info(fmt.Sprintf("vp8tunnel: peer restart detected, new epoch=0x%08x", res.PeerEpoch))
if t.OnPeerRestart != nil {
t.OnPeerRestart()
}
}
if res.Keepalive || len(res.Payload) == 0 {
return
}
n := t.recvFrames.Add(1)
if n <= 5 || n%500 == 0 {
t.logger.Debug(fmt.Sprintf("vp8tunnel: recv frame #%d size=%d", n, len(res.Payload)))
}
if t.OnData != nil {
t.OnData(res.Payload)
}
}
func (t *VP8DataTunnel) currentRate() (fps, batch int) {
t.cfgMu.Lock()
defer t.cfgMu.Unlock()
return t.fps, t.batch
}
func sampleIntervalFor(fps, batch int) time.Duration {
if fps < 1 {
fps = 1
}
frameInterval := time.Second / time.Duration(fps)
interval := frameInterval
if batch > 1 {
interval = frameInterval / time.Duration(batch)
}
if interval <= 0 {
interval = time.Millisecond
}
return interval
}
func pacedBatchFor(batch int) int {
if batch <= 1 {
return batch
}
floor := batch * paceBatchFloorPercent / 100
if floor < 1 {
floor = 1
}
return common.IntInRange(floor, batch)
}
func (t *VP8DataTunnel) writerLoop() {
for {
fps, batch := t.currentRate()
pacedBatch := pacedBatchFor(batch)
sampleInterval := sampleIntervalFor(fps, pacedBatch)
keepaliveEvery, keepalivePad := t.nextKeepalive(sampleInterval)
t.logger.Debug(fmt.Sprintf("vp8tunnel: writer (re)started fps=%d batch=%d pacedBatch=%d sampleInterval=%s keepaliveEvery=%d",
fps, batch, pacedBatch, sampleInterval, keepaliveEvery))
ticker := time.NewTicker(sampleInterval)
drift := time.NewTimer(common.DurationInRange(paceDriftMin, paceDriftMax))
idleTicks := 0
reconfigure := false
for !reconfigure {
select {
case <-t.stopCh:
ticker.Stop()
drift.Stop()
return
case <-t.cfgChan:
reconfigure = true
case <-drift.C:
pacedBatch = pacedBatchFor(batch)
sampleInterval = sampleIntervalFor(fps, pacedBatch)
ticker.Reset(sampleInterval)
keepaliveEvery, keepalivePad = t.nextKeepalive(sampleInterval)
drift.Reset(common.DurationInRange(paceDriftMin, paceDriftMax))
t.logger.Debug(fmt.Sprintf("vp8tunnel: pace drift pacedBatch=%d/%d sampleInterval=%s", pacedBatch, batch, sampleInterval))
case <-ticker.C:
var sample []byte
isKeepalive := false
select {
case data := <-t.sendQueue:
sample = t.obf.EncodeData(data)
idleTicks = 0
default:
idleTicks++
if idleTicks < keepaliveEvery {
continue
}
idleTicks = 0
sample = t.obf.EncodeKeepalive(keepalivePad)
keepaliveEvery, keepalivePad = t.nextKeepalive(sampleInterval)
isKeepalive = true
}
if sample == nil {
continue
}
if err := t.track.WriteSample(media.Sample{Data: sample, Duration: sampleInterval}); err != nil {
t.logger.Debug(fmt.Sprintf("vp8tunnel: WriteSample error: %v", err))
continue
}
n := t.sentFrames.Add(1)
if isKeepalive {
t.keepaliveFrames.Add(1)
}
if n <= 5 || n%500 == 0 {
keepalives := t.keepaliveFrames.Load()
t.logger.Debug(fmt.Sprintf("vp8tunnel: sent frame #%d size=%d data=%d keepalive=%d", n, len(sample), n-keepalives, keepalives))
}
}
}
ticker.Stop()
drift.Stop()
}
}

284
transport/call/vk/api.go Normal file
View File

@@ -0,0 +1,284 @@
package vk
import (
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
type TurnServer struct {
URLs []string `json:"urls"`
Username string `json:"username"`
Credential string `json:"credential"`
}
type StunServer struct {
URLs []string `json:"urls"`
}
type CallInfo struct {
CallID string
JoinLink string
ShortLink string
OKJoinLink string
TurnServer TurnServer
StunServer StunServer
WSEndpoint string
}
type vkTokenResponse struct {
Data struct {
AccessToken string `json:"access_token"`
} `json:"data"`
}
type callSettingsResponse struct {
Response struct {
Settings struct {
PublicKey string `json:"public_key"`
} `json:"settings"`
} `json:"response"`
}
type callTokenResponse struct {
Response struct {
Token string `json:"token"`
APIBaseURL string `json:"api_base_url"`
} `json:"response"`
}
type okAuthResponse struct {
SessionKey string `json:"session_key"`
}
type joinResponse struct {
Endpoint string `json:"endpoint"`
TurnServer TurnServer `json:"turn_server"`
StunServer StunServer `json:"stun_server"`
}
func JoinExistingCall(dialer N.Dialer, cookieStr, vkLink string, cfg VKConfig, logger logger.ContextLogger) (*CallInfo, error) {
if cfg.AppID == "" || cfg.APIVersion == "" {
return nil, fmt.Errorf("config incomplete: app_id=%q api=%q", cfg.AppID, cfg.APIVersion)
}
token := extractJoinToken(vkLink)
if token == "" {
return nil, fmt.Errorf("could not extract join token from %q", vkLink)
}
logger.Info(fmt.Sprintf("[auth] Joining existing call token=%s", token))
resp, err := authAndJoin(dialer, cookieStr, token, cfg)
if err != nil {
return nil, err
}
return &CallInfo{
JoinLink: vkLink,
OKJoinLink: token,
TurnServer: resp.TurnServer,
StunServer: resp.StunServer,
WSEndpoint: resp.Endpoint,
}, nil
}
func CreateAndJoinCall(dialer N.Dialer, cookieStr, peerId string, cfg VKConfig, logger logger.ContextLogger) (*CallInfo, error) {
if cfg.AppID == "" || cfg.APIVersion == "" {
return nil, fmt.Errorf("config incomplete: app_id=%q api=%q", cfg.AppID, cfg.APIVersion)
}
auth := func(bearer string) map[string]string {
return map[string]string{"Authorization": "Bearer " + bearer}
}
logger.Info("[auth] Getting VK token...")
r, err := httpPost(dialer, "https://login.vk.com/?act=web_token",
url.Values{"version": {"1"}, "app_id": {cfg.AppID}},
map[string]string{"Cookie": cookieStr})
if err != nil {
return nil, fmt.Errorf("web_token: %w", err)
}
var tok vkTokenResponse
json.Unmarshal(r, &tok)
vkToken := tok.Data.AccessToken
if vkToken == "" {
return nil, fmt.Errorf("empty VK token, response: %s", string(r))
}
logger.Info(fmt.Sprintf("[auth] Creating call peer_id=%s...", peerId))
r, err = httpPost(dialer, "https://api.vk.com/method/calls.start",
url.Values{"v": {cfg.APIVersion}, "peer_id": {peerId}}, auth(vkToken))
if err != nil {
return nil, fmt.Errorf("calls.start: %w", err)
}
var call struct {
Response struct {
CallID string `json:"call_id"`
JoinLink string `json:"join_link"`
OKJoinLink string `json:"ok_join_link"`
ShortCredentials struct {
LinkWithPassword string `json:"link_with_password"`
} `json:"short_credentials"`
} `json:"response"`
}
json.Unmarshal(r, &call)
c := call.Response
if c.CallID == "" {
return nil, fmt.Errorf("empty call_id, response: %s", string(r))
}
if c.OKJoinLink == "" {
return nil, fmt.Errorf("empty ok_join_link, response: %s", string(r))
}
logger.Debug(fmt.Sprintf("[auth] call_id: %s", c.CallID))
logger.Debug(fmt.Sprintf("[auth] join_link: %s", c.JoinLink))
logger.Info("[auth] Joining conversation...")
resp, err := authAndJoin(dialer, cookieStr, c.OKJoinLink, cfg)
if err != nil {
return nil, err
}
return &CallInfo{
CallID: c.CallID, JoinLink: c.JoinLink, ShortLink: c.ShortCredentials.LinkWithPassword,
OKJoinLink: c.OKJoinLink, TurnServer: resp.TurnServer, StunServer: resp.StunServer,
WSEndpoint: resp.Endpoint,
}, nil
}
func BuildICEServers(callInfo *CallInfo) []ICEServerSpec {
var servers []ICEServerSpec
if len(callInfo.StunServer.URLs) > 0 {
servers = append(servers, ICEServerSpec{URLs: callInfo.StunServer.URLs})
}
if len(callInfo.TurnServer.URLs) > 0 {
urls := append([]string{}, callInfo.TurnServer.URLs...)
urls = append(urls, urls[len(urls)-1]+"?transport=tcp")
servers = append(servers, ICEServerSpec{
URLs: urls, Username: callInfo.TurnServer.Username, Credential: callInfo.TurnServer.Credential,
})
}
return servers
}
type ICEServerSpec struct {
URLs []string
Username string
Credential string
}
func authAndJoin(dialer N.Dialer, cookieStr, okJoinLink string, cfg VKConfig) (*joinResponse, error) {
auth := func(bearer string) map[string]string {
return map[string]string{"Authorization": "Bearer " + bearer}
}
r, err := httpPost(dialer, "https://login.vk.com/?act=web_token",
url.Values{"version": {"1"}, "app_id": {cfg.AppID}},
map[string]string{"Cookie": cookieStr})
if err != nil {
return nil, fmt.Errorf("web_token: %w", err)
}
var tok vkTokenResponse
json.Unmarshal(r, &tok)
if tok.Data.AccessToken == "" {
return nil, fmt.Errorf("empty VK token, response: %s", string(r))
}
r, err = httpPost(dialer, "https://api.vk.com/method/calls.getSettings",
url.Values{"v": {cfg.APIVersion}}, auth(tok.Data.AccessToken))
if err != nil {
return nil, fmt.Errorf("calls.getSettings: %w", err)
}
var settings callSettingsResponse
json.Unmarshal(r, &settings)
appKey := settings.Response.Settings.PublicKey
if appKey == "" {
return nil, fmt.Errorf("empty public_key, response: %s", string(r))
}
r, err = httpPost(dialer, "https://api.vk.com/method/messages.getCallToken",
url.Values{"v": {cfg.APIVersion}, "env": {"production"}}, auth(tok.Data.AccessToken))
if err != nil {
return nil, fmt.Errorf("messages.getCallToken: %w", err)
}
var callToken callTokenResponse
json.Unmarshal(r, &callToken)
if callToken.Response.Token == "" {
return nil, fmt.Errorf("empty call token, response: %s", string(r))
}
if callToken.Response.APIBaseURL == "" {
return nil, fmt.Errorf("empty api_base_url, response: %s", string(r))
}
apiBaseURL := strings.TrimRight(callToken.Response.APIBaseURL, "/")
if !strings.HasSuffix(apiBaseURL, "/fb.do") {
apiBaseURL += "/fb.do"
}
sd, _ := json.Marshal(map[string]interface{}{
"device_id": "sing-box-go-1", "client_version": cfg.AppVersion,
"client_type": "SDK_JS", "auth_token": callToken.Response.Token, "version": 3,
})
r, err = httpPost(dialer, apiBaseURL, url.Values{
"method": {"auth.anonymLogin"}, "application_key": {appKey},
"format": {"json"}, "session_data": {string(sd)},
}, nil)
if err != nil {
return nil, fmt.Errorf("auth.anonymLogin: %w", err)
}
var okAuth okAuthResponse
json.Unmarshal(r, &okAuth)
if okAuth.SessionKey == "" {
return nil, fmt.Errorf("empty session_key, response: %s", string(r))
}
ms, _ := json.Marshal(map[string]bool{
"isAudioEnabled": false, "isVideoEnabled": true, "isScreenSharingEnabled": false,
})
r, err = httpPost(dialer, apiBaseURL, url.Values{
"method": {"vchat.joinConversationByLink"}, "session_key": {okAuth.SessionKey},
"application_key": {appKey}, "format": {"json"}, "joinLink": {okJoinLink},
"isVideo": {"true"}, "isAudio": {"false"}, "mediaSettings": {string(ms)},
}, nil)
if err != nil {
return nil, fmt.Errorf("vchat.joinConversationByLink: %w", err)
}
var jr joinResponse
json.Unmarshal(r, &jr)
if jr.Endpoint == "" {
return nil, fmt.Errorf("empty WS endpoint, response: %s", string(r))
}
return &jr, nil
}
func extractJoinToken(link string) string {
link = strings.TrimSpace(link)
if link == "" {
return ""
}
if u, err := url.Parse(link); err == nil && u.Scheme != "" {
path := strings.Trim(u.Path, "/")
if path != "" {
parts := strings.Split(path, "/")
return parts[len(parts)-1]
}
}
if !strings.ContainsAny(link, "/?&=") {
return link
}
parts := strings.Split(strings.TrimRight(link, "/"), "/")
return parts[len(parts)-1]
}
func httpPost(dialer N.Dialer, endpoint string, form url.Values, extraHeaders map[string]string) ([]byte, error) {
body := form.Encode()
req, err := http.NewRequest("POST", endpoint, strings.NewReader(body))
if err != nil {
return nil, err
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("User-Agent", common.UserAgent)
req.Header.Set("Origin", "https://vk.com")
req.Header.Set("Referer", "https://vk.com/")
for k, v := range extraHeaders {
req.Header.Set(k, v)
}
resp, err := common.HttpClient(dialer).Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
return io.ReadAll(resp.Body)
}

View File

@@ -0,0 +1,375 @@
package vk
import (
"bytes"
"compress/gzip"
"context"
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"net/http/httputil"
"net/url"
"strings"
"sync"
"time"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
var activeCaptchaProxy struct {
sync.Mutex
listener net.Listener
port int
keyCh chan string
doneCh chan struct{}
}
func StartCaptchaProxy(redirectURI string, dialer N.Dialer) int {
StopCaptchaProxy()
targetURL, err := url.Parse(redirectURI)
if err != nil {
return 0
}
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
return 0
}
port := listener.Addr().(*net.TCPAddr).Port
localOrigin := fmt.Sprintf("http://127.0.0.1:%d", port)
upstreamOrigin := targetURL.Scheme + "://" + targetURL.Host
keyCh := make(chan string, 1)
transport := &http.Transport{
MaxIdleConns: 100,
MaxIdleConnsPerHost: 100,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ForceAttemptHTTP2: false,
}
if dialer != nil {
transport.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
return dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
}
}
proxy := &httputil.ReverseProxy{
Transport: transport,
Rewrite: func(req *httputil.ProxyRequest) {
req.Out.URL.Scheme = targetURL.Scheme
req.Out.URL.Host = targetURL.Host
if req.Out.URL.Path == "" {
req.Out.URL.Path = targetURL.Path
}
req.Out.Host = targetURL.Host
req.Out.Header.Del("Accept-Encoding")
req.Out.Header.Del("TE")
for _, headerName := range []string{"Origin", "Referer"} {
val := req.Out.Header.Get(headerName)
if val != "" {
req.Out.Header.Set(headerName, strings.ReplaceAll(val, localOrigin, upstreamOrigin))
}
}
},
ModifyResponse: func(res *http.Response) error {
rewriteProxyCookies(res)
if res.StatusCode >= 300 && res.StatusCode < 400 {
if loc := res.Header.Get("Location"); loc != "" {
res.Header.Set("Location", strings.ReplaceAll(loc, upstreamOrigin, localOrigin))
}
}
contentType := res.Header.Get("Content-Type")
shouldInspect := isHTMLLike(contentType) || strings.Contains(res.Request.URL.Path, "captchaNotRobot.check")
if !shouldInspect {
return nil
}
reader := res.Body
decompressed := false
if res.Header.Get("Content-Encoding") == "gzip" {
gzReader, err := gzip.NewReader(res.Body)
if err == nil {
reader = gzReader
decompressed = true
defer gzReader.Close()
}
}
bodyBytes, err := io.ReadAll(reader)
if err != nil {
return err
}
res.Body.Close()
if strings.Contains(res.Request.URL.Path, "captchaNotRobot.check") {
token := extractSuccessToken(bodyBytes)
if token != "" {
select {
case keyCh <- token:
default:
}
}
}
if isHTMLLike(contentType) {
for _, h := range []string{
"Content-Security-Policy", "Content-Security-Policy-Report-Only",
"X-Content-Security-Policy", "X-WebKit-CSP",
"Cross-Origin-Opener-Policy", "Cross-Origin-Embedder-Policy",
"Cross-Origin-Resource-Policy", "X-Frame-Options",
"Strict-Transport-Security", "Alt-Svc",
} {
res.Header.Del(h)
}
bodyBytes = []byte(rewriteCaptchaHTML(string(bodyBytes), localOrigin, upstreamOrigin))
}
if decompressed {
res.Header.Del("Content-Encoding")
}
res.Body = io.NopCloser(bytes.NewReader(bodyBytes))
res.ContentLength = int64(len(bodyBytes))
res.Header.Set("Content-Length", fmt.Sprint(len(bodyBytes)))
return nil
},
}
mux := http.NewServeMux()
mux.HandleFunc("/local-captcha-result", func(w http.ResponseWriter, r *http.Request) {
token := r.FormValue("token")
if token != "" {
select {
case keyCh <- token:
default:
}
}
w.Header().Set("Access-Control-Allow-Origin", "*")
fmt.Fprint(w, "ok")
})
mux.HandleFunc("/generic_proxy", func(w http.ResponseWriter, r *http.Request) {
proxyURL := r.URL.Query().Get("proxy_url")
parsed, err := url.Parse(proxyURL)
if err != nil || parsed.Host == "" {
http.Error(w, "Bad URL", http.StatusBadRequest)
return
}
genericProxy := &httputil.ReverseProxy{
Transport: transport,
Rewrite: func(req *httputil.ProxyRequest) {
req.Out.URL.Scheme = parsed.Scheme
req.Out.URL.Host = parsed.Host
req.Out.URL.Path = parsed.Path
req.Out.URL.RawQuery = parsed.RawQuery
req.Out.Host = parsed.Host
req.Out.Header.Del("Accept-Encoding")
},
}
genericProxy.ServeHTTP(w, r)
})
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/" && targetURL.Path != "" && targetURL.Path != "/" && r.URL.RawQuery == "" {
localPath := targetURL.Path
if targetURL.RawQuery != "" {
localPath += "?" + targetURL.RawQuery
}
http.Redirect(w, r, localPath, http.StatusTemporaryRedirect)
return
}
proxy.ServeHTTP(w, r)
})
activeCaptchaProxy.Lock()
activeCaptchaProxy.listener = listener
activeCaptchaProxy.port = port
activeCaptchaProxy.keyCh = keyCh
activeCaptchaProxy.doneCh = make(chan struct{})
activeCaptchaProxy.Unlock()
go http.Serve(listener, mux)
return port
}
func GetCaptchaResult() string {
activeCaptchaProxy.Lock()
ch := activeCaptchaProxy.keyCh
done := activeCaptchaProxy.doneCh
activeCaptchaProxy.Unlock()
if ch == nil || done == nil {
return ""
}
select {
case token := <-ch:
return token
case <-done:
return ""
case <-time.After(300 * time.Second):
return ""
}
}
func StopCaptchaProxy() {
activeCaptchaProxy.Lock()
ln := activeCaptchaProxy.listener
done := activeCaptchaProxy.doneCh
activeCaptchaProxy.listener = nil
activeCaptchaProxy.port = 0
activeCaptchaProxy.keyCh = nil
activeCaptchaProxy.doneCh = nil
activeCaptchaProxy.Unlock()
if done != nil {
close(done)
}
if ln != nil {
ln.Close()
}
}
func rewriteProxyCookies(res *http.Response) {
cookies := res.Cookies()
if len(cookies) == 0 {
return
}
res.Header.Del("Set-Cookie")
for _, cookie := range cookies {
cookie.Domain = ""
cookie.Secure = false
cookie.Partitioned = false
if cookie.SameSite == http.SameSiteNoneMode || cookie.SameSite == http.SameSiteStrictMode {
cookie.SameSite = http.SameSiteLaxMode
}
res.Header.Add("Set-Cookie", cookie.String())
}
}
func isHTMLLike(contentType string) bool {
return strings.Contains(contentType, "text/html") ||
strings.Contains(contentType, "application/xhtml+xml")
}
func extractSuccessToken(body []byte) string {
var payload struct {
Response struct {
SuccessToken string `json:"success_token"`
} `json:"response"`
}
if err := json.Unmarshal(body, &payload); err != nil {
return ""
}
return payload.Response.SuccessToken
}
func rewriteCaptchaHTML(html, localOrigin, upstreamOrigin string) string {
html = strings.ReplaceAll(html, upstreamOrigin, localOrigin)
script := fmt.Sprintf(`
<script>
(function() {
var localOrigin = %q;
var upstreamOrigin = %q;
function rewriteUrl(urlStr) {
if (!urlStr || typeof urlStr !== 'string') return urlStr;
if (urlStr.indexOf(localOrigin) === 0) return urlStr;
if (urlStr.indexOf(upstreamOrigin) === 0) return localOrigin + urlStr.slice(upstreamOrigin.length);
if (urlStr.indexOf('//') === 0) {
return '/generic_proxy?proxy_url=' + encodeURIComponent(window.location.protocol + urlStr);
}
if (urlStr.indexOf('http://') === 0 || urlStr.indexOf('https://') === 0) {
return '/generic_proxy?proxy_url=' + encodeURIComponent(urlStr);
}
return urlStr;
}
function rewriteElementAttr(el, attr) {
if (!el || !el.getAttribute) return;
var value = el.getAttribute(attr);
if (!value) return;
var rewritten = rewriteUrl(value);
if (rewritten !== value) el.setAttribute(attr, rewritten);
}
function rewriteDocument(root) {
if (!root || !root.querySelectorAll) return;
root.querySelectorAll('[href]').forEach(function(el) { rewriteElementAttr(el, 'href'); });
root.querySelectorAll('[src]').forEach(function(el) { rewriteElementAttr(el, 'src'); });
root.querySelectorAll('form[action]').forEach(function(el) { rewriteElementAttr(el, 'action'); });
}
function handleSuccessToken(token) {
if (!token) return;
fetch('/local-captcha-result', {
method: 'POST',
headers: {'Content-Type': 'application/x-www-form-urlencoded'},
body: 'token=' + encodeURIComponent(token)
}).catch(function() {});
}
var origOpen = XMLHttpRequest.prototype.open;
XMLHttpRequest.prototype.open = function() {
if (arguments[1] && typeof arguments[1] === 'string') {
this._origUrl = arguments[1];
arguments[1] = rewriteUrl(arguments[1]);
}
return origOpen.apply(this, arguments);
};
var origSend = XMLHttpRequest.prototype.send;
XMLHttpRequest.prototype.send = function() {
var xhr = this;
if (this._origUrl && this._origUrl.indexOf('captchaNotRobot.check') !== -1) {
xhr.addEventListener('load', function() {
try {
var data = JSON.parse(xhr.responseText);
if (data.response && data.response.success_token) handleSuccessToken(data.response.success_token);
} catch (e) {}
});
}
return origSend.apply(this, arguments);
};
var origFetch = window.fetch;
if (origFetch) {
window.fetch = function() {
var url = arguments[0];
var urlStr = (typeof url === 'object' && url && url.url) ? url.url : url;
var origUrlStr = urlStr;
if (typeof urlStr === 'string') {
urlStr = rewriteUrl(urlStr);
arguments[0] = urlStr;
}
var p = origFetch.apply(this, arguments);
if (typeof origUrlStr === 'string' && origUrlStr.indexOf('captchaNotRobot.check') !== -1) {
p.then(function(r) { return r.clone().json(); }).then(function(data) {
if (data.response && data.response.success_token) handleSuccessToken(data.response.success_token);
}).catch(function() {});
}
return p;
};
}
var origWindowOpen = window.open;
if (origWindowOpen) {
window.open = function(url) {
if (typeof url === 'string') arguments[0] = rewriteUrl(url);
return origWindowOpen.apply(this, arguments);
};
}
rewriteDocument(document);
if (document.documentElement && window.MutationObserver) {
new MutationObserver(function(mutations) {
mutations.forEach(function(mutation) {
if (mutation.type === 'attributes' && mutation.target) {
rewriteElementAttr(mutation.target, mutation.attributeName);
return;
}
mutation.addedNodes.forEach(function(node) {
if (node.nodeType === 1) rewriteDocument(node);
});
});
}).observe(document.documentElement, {
subtree: true, childList: true, attributes: true,
attributeFilter: ['href', 'src', 'action']
});
}
})();
</script>
`, localOrigin, upstreamOrigin)
if idx := strings.Index(html, "</head>"); idx >= 0 {
return html[:idx] + script + html[idx:]
}
if idx := strings.Index(html, "</body>"); idx >= 0 {
return html[:idx] + script + html[idx:]
}
return html + script
}

View File

@@ -0,0 +1,28 @@
package vk
import (
"fmt"
"github.com/sagernet/sing/common/logger"
)
type VKConfig struct {
AppID string
APIVersion string
SDKVersion string
AppVersion string
ProtocolVersion string
}
func FetchConfig(logger logger.ContextLogger) (VKConfig, error) {
cfg := VKConfig{
AppID: "6287487",
APIVersion: "5.280",
SDKVersion: "2.8.6-beta.22",
AppVersion: "1.1",
ProtocolVersion: "6",
}
logger.Debug(fmt.Sprintf("[config] app_id=%s api=%s sdk=%s app=%s proto=%s",
cfg.AppID, cfg.APIVersion, cfg.SDKVersion, cfg.AppVersion, cfg.ProtocolVersion))
return cfg, nil
}

View File

@@ -0,0 +1,125 @@
package vk
import (
"context"
"encoding/json"
"fmt"
"time"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing-box/transport/call/tunnel"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
func ConnectCreator(ctx context.Context, cookieStr, joinLink string, readBuf int, dialer N.Dialer, logger logger.ContextLogger) (*tunnel.RelayBridge, string, error) {
cfg, err := FetchConfig(logger)
if err != nil {
return nil, "", err
}
var callInfo *CallInfo
if joinLink != "" {
callInfo, err = JoinExistingCall(dialer, cookieStr, joinLink, cfg, logger)
} else {
callInfo, err = CreateAndJoinCall(dialer, cookieStr, "", cfg, logger)
}
if err != nil {
return nil, "", err
}
if readBuf <= 0 {
readBuf = 32768
}
obf, err := tunnel.NewTunnelObfuscator(tunnel.DeriveSecretFromJoinLink(callInfo.JoinLink))
if err != nil {
return nil, "", fmt.Errorf("vk: obfuscator init: %w", err)
}
bridge := &Bridge{
dialer: dialer,
readBuf: readBuf,
logger: logger,
}
bridge.newRelay = func() Relay {
ur := NewTunnelRelay(dialer, logger)
ur.readBufSize = readBuf
ur.SetObfuscator(obf)
ur.OnConnected = func(tun tunnel.DataTunnel) {
bridgeReadBuf := common.VP8BufSize
if _, ok := tun.(*tunnel.DCTunnel); ok {
bridgeReadBuf = readBuf
}
rb := tunnel.NewRelayBridge(tun, "creator", bridgeReadBuf, dialer, logger)
rb.MarkReady()
if st, ok := tun.(*tunnel.SymmetricScreenTunnel); ok {
rb.SetOnPeerConfig(func(fps, batch, trackCount int) {
st.SetTrackCount(trackCount)
bridge.setScreenSharing(trackCount > 1)
})
}
bridge.mu.Lock()
bridge.activeBridge = rb
bridge.mu.Unlock()
}
return ur
}
go bridge.Run(callInfo, cookieStr, cfg)
deadline := time.Now().Add(60 * time.Second)
for {
bridge.mu.Lock()
rb := bridge.activeBridge
bridge.mu.Unlock()
if rb != nil {
return rb, callInfo.JoinLink, nil
}
if time.Now().After(deadline) {
return nil, "", fmt.Errorf("vk: creator tunnel timed out")
}
select {
case <-ctx.Done():
return nil, "", ctx.Err()
case <-time.After(200 * time.Millisecond):
}
}
}
func ConnectJoiner(ctx context.Context, joinLink, displayName string, readBuf int, dialer N.Dialer, dnsRouter adapter.DNSRouter, logger logger.ContextLogger) (tunnel.DataTunnel, error) {
if displayName == "" {
displayName = "Joiner"
}
authJSON, err := RunVKAuth(dialer, joinLink, displayName, logger)
if err != nil {
return nil, fmt.Errorf("vk: auth: %w", err)
}
var params VKAuthParams
if err := json.Unmarshal([]byte(authJSON), &params); err != nil {
return nil, fmt.Errorf("vk: decode auth params: %w", err)
}
params.TunnelMode = "video"
paramsJSON, err := json.Marshal(params)
if err != nil {
return nil, fmt.Errorf("vk: encode auth params: %w", err)
}
joiner := NewVKJoiner(
logger,
nil,
common.AddTunnelTracks,
common.ReadTrack,
dialer,
dnsRouter,
)
tunCh := make(chan tunnel.DataTunnel, 1)
joiner.OnConnected = func(tun tunnel.DataTunnel) {
select {
case tunCh <- tun:
default:
}
}
go joiner.RunWithParams(string(paramsJSON))
select {
case tun := <-tunCh:
return tun, nil
case <-ctx.Done():
joiner.Close()
return nil, ctx.Err()
}
}

View File

@@ -0,0 +1,318 @@
package vk
import (
"context"
"encoding/json"
"fmt"
"net"
"net/http"
"strings"
"sync"
"time"
"github.com/gorilla/websocket"
"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"
)
const topologyDirect = "DIRECT"
const maxServerBounces = 5
type Bridge struct {
mu sync.Mutex
vkWs *websocket.Conn
vkSeq int
iceServers []webrtc.ICEServer
topology string
peers map[int64]struct{}
relay Relay
newRelay func() Relay
p2p *P2PHandler
screenSharing bool
serverBounces int
suppressScreenshare bool
bouncing bool
dialer N.Dialer
activeBridge *tunnel.RelayBridge
readBuf int
logger logger.ContextLogger
}
func (b *Bridge) setScreenSharing(enabled bool) {
b.mu.Lock()
if b.vkWs == nil || b.screenSharing == enabled {
b.mu.Unlock()
return
}
if enabled && b.suppressScreenshare {
b.mu.Unlock()
b.logger.Debug("[vk-ws] screenshare suppressed after SERVER flap, staying single-track DIRECT")
return
}
b.screenSharing = enabled
b.mu.Unlock()
b.logger.Debug(fmt.Sprintf("[vk-ws] peer track count change, screenshare=%v", enabled))
b.sendMediaSettings(enabled)
}
func (b *Bridge) vkSend(command string, extra map[string]interface{}) {
b.mu.Lock()
defer b.mu.Unlock()
if b.vkWs == nil {
return
}
b.vkSeq++
seq := b.vkSeq
var out []byte
if pid, ok := extra["participantId"]; ok {
dataJSON, _ := json.Marshal(extra["data"])
out = []byte(fmt.Sprintf(`{"command":%q,"sequence":%d,"participantId":%v,"data":%s}`,
command, seq, pid, dataJSON))
} else {
extra["command"] = command
extra["sequence"] = seq
out, _ = json.Marshal(extra)
}
b.vkWs.WriteMessage(websocket.TextMessage, out)
b.logger.Debug(fmt.Sprintf("[vk-ws] -> %s", command))
}
func (b *Bridge) sendMediaSettings(screenSharing bool) {
b.vkSend("change-media-settings", map[string]interface{}{
"mediaSettings": map[string]interface{}{
"isAudioEnabled": false, "isVideoEnabled": true,
"isScreenSharingEnabled": screenSharing, "isFastScreenSharingEnabled": false,
"isAudioSharingEnabled": false, "isAnimojiEnabled": false,
},
})
}
func (b *Bridge) handleVKMessage(raw []byte) {
var msg map[string]interface{}
if err := json.Unmarshal(raw, &msg); err != nil {
return
}
msgType, _ := msg["type"].(string)
switch msgType {
case "notification":
notif, _ := msg["notification"].(string)
b.logger.Debug(fmt.Sprintf("[vk-ws] <- notification: %s", notif))
switch notif {
case "connection":
b.logger.Debug("[vk-ws] TURN creds received")
case "transmitted-data":
data, _ := msg["data"].(map[string]interface{})
if data != nil && b.topology == topologyDirect && b.p2p != nil {
b.p2p.OnTransmittedData(data)
}
case "registered-peer":
pid, _ := msg["participantId"].(float64)
if b.topology == topologyDirect && b.p2p != nil {
b.p2p.OnRegisteredPeer(int64(pid))
}
case "topology-changed":
topo, _ := msg["topology"].(string)
b.logger.Debug(fmt.Sprintf("[vk-ws] Topology changed to %s", topo))
b.topology = topo
if topo != topologyDirect {
b.bounceForServerTopology("SERVER topology")
return
}
case "participant-joined", "participant-added":
if pid, ok := msg["participantId"].(float64); ok {
b.peers[int64(pid)] = struct{}{}
b.logger.Debug(fmt.Sprintf("[vk-ws] Participant %d joined (total: %d)", int64(pid), len(b.peers)))
if b.topology != topologyDirect {
b.bounceForServerTopology("participant joined under SERVER")
return
}
}
case "participant-left":
if pid, ok := msg["participantId"].(float64); ok {
delete(b.peers, int64(pid))
b.logger.Debug(fmt.Sprintf("[vk-ws] Participant %d left (total: %d)", int64(pid), len(b.peers)))
}
case "hungup":
if pid, ok := msg["participantId"].(float64); ok {
delete(b.peers, int64(pid))
b.logger.Debug(fmt.Sprintf("[vk-ws] Participant %d hung up (total: %d)", int64(pid), len(b.peers)))
} else {
b.logger.Debug("[vk-ws] Participant hung up")
}
case "closed-conversation":
reason, _ := msg["reason"].(string)
b.logger.Debug(fmt.Sprintf("[vk-ws] Conversation closed: %s", reason))
b.mu.Lock()
if b.vkWs != nil {
b.vkWs.Close()
}
b.mu.Unlock()
default:
snippet, _ := json.Marshal(msg)
if len(snippet) > 1000 {
snippet = append(snippet[:1000], '.', '.', '.')
}
b.logger.Debug(fmt.Sprintf("[vk-ws] unhandled: %s", string(snippet)))
}
case "response":
seq, _ := msg["sequence"].(float64)
snippet, _ := json.Marshal(msg)
if len(snippet) > 1000 {
snippet = append(snippet[:1000], '.', '.', '.')
}
b.logger.Debug(fmt.Sprintf("[vk-ws] <- response seq=%d: %s", int(seq), string(snippet)))
case "error":
errMsg, _ := msg["message"].(string)
errCode, _ := msg["error"].(string)
b.logger.Warn(fmt.Sprintf("[vk-ws] <- error: %s %s", errCode, errMsg))
}
}
func (b *Bridge) connectVKWs(wsURL string) error {
vkHeader := http.Header{}
vkHeader.Set("User-Agent", common.UserAgent)
vkHeader.Set("Origin", "https://vk.com")
vkDialer := websocket.Dialer{
WriteBufferSize: common.RTPBufSize,
NetDialContext: b.dialContext,
}
vkWs, _, err := vkDialer.Dial(wsURL, vkHeader)
if err != nil {
return err
}
b.mu.Lock()
b.vkWs = vkWs
b.vkSeq = 0
b.mu.Unlock()
return nil
}
func (b *Bridge) dialContext(ctx context.Context, network, addr string) (net.Conn, error) {
return b.dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
}
func (b *Bridge) initRelay() {
if b.relay != nil {
b.relay.Close()
}
b.topology = topologyDirect
b.peers = make(map[int64]struct{})
b.relay = b.newRelay()
b.p2p = NewP2PHandler(b)
b.p2p.Init()
}
func (b *Bridge) bounceForServerTopology(reason string) {
b.mu.Lock()
if b.bouncing {
b.mu.Unlock()
return
}
b.bouncing = true
b.serverBounces++
count := b.serverBounces
if count > maxServerBounces {
b.suppressScreenshare = true
}
suppress := b.suppressScreenshare
ws := b.vkWs
b.mu.Unlock()
if suppress {
b.logger.Debug(fmt.Sprintf("[vk-ws] %s -> reconnect #%d, suppressing screenshare to settle single-track DIRECT", reason, count))
} else {
b.logger.Debug(fmt.Sprintf("[vk-ws] %s -> manual reconnect #%d to recover DIRECT", reason, count))
}
if ws != nil {
ws.Close()
}
}
func (b *Bridge) readLoop() error {
for {
_, msg, err := b.vkWs.ReadMessage()
if err != nil {
return err
}
if string(msg) == "ping" {
b.mu.Lock()
b.vkWs.WriteMessage(websocket.TextMessage, []byte("pong"))
b.mu.Unlock()
continue
}
b.handleVKMessage(msg)
}
}
func (b *Bridge) Run(callInfo *CallInfo, cookieStr string, cfg VKConfig) {
b.logger.Info(fmt.Sprintf("CALL CREATED join_link=%s turn=%s protocol=v%s sdk=%s",
callInfo.JoinLink, strings.Join(callInfo.TurnServer.URLs, ", "), cfg.ProtocolVersion, cfg.SDKVersion))
b.iceServers = buildWebRTCICEServers(BuildICEServers(callInfo))
wsEndpoint := callInfo.WSEndpoint
capabilities := "2F7F"
makeWSURL := func(ep string) string {
return ep +
"&platform=WEB" +
"&appVersion=" + cfg.AppVersion +
"&version=" + cfg.ProtocolVersion +
"&device=browser&capabilities=" + capabilities + "&clientType=VK&tgt=join"
}
go func() {
for {
time.Sleep(15 * time.Second)
b.mu.Lock()
ws := b.vkWs
b.mu.Unlock()
if ws != nil {
b.mu.Lock()
ws.WriteMessage(websocket.PingMessage, nil)
b.mu.Unlock()
}
}
}()
for {
b.initRelay()
b.logger.Debug("[vk-ws] Connecting...")
if err := b.connectVKWs(makeWSURL(wsEndpoint)); err != nil {
b.logger.Warn(fmt.Sprintf("[vk-ws] Connect failed: %s, retrying in 5s...", common.MaskError(err)))
time.Sleep(5 * time.Second)
continue
}
b.logger.Debug("[vk-ws] Connected")
b.mu.Lock()
b.screenSharing = false
b.bouncing = false
b.mu.Unlock()
b.sendMediaSettings(false)
err := b.readLoop()
b.logger.Debug(fmt.Sprintf("[vk-ws] Closed: %s", common.MaskError(err)))
b.mu.Lock()
b.vkWs = nil
b.mu.Unlock()
b.logger.Debug("[vk-ws] Rejoining in 3s...")
time.Sleep(3 * time.Second)
joinResp, rerr := authAndJoin(b.dialer, cookieStr, callInfo.OKJoinLink, cfg)
if rerr != nil {
b.logger.Warn(fmt.Sprintf("[rejoin] Failed: %v, retrying in 5s...", rerr))
time.Sleep(5 * time.Second)
continue
}
wsEndpoint = joinResp.Endpoint
callInfo.TurnServer = joinResp.TurnServer
callInfo.StunServer = joinResp.StunServer
b.iceServers = buildWebRTCICEServers(BuildICEServers(callInfo))
}
}
func buildWebRTCICEServers(specs []ICEServerSpec) []webrtc.ICEServer {
out := make([]webrtc.ICEServer, len(specs))
for i, s := range specs {
out[i] = webrtc.ICEServer{URLs: s.URLs, Username: s.Username, Credential: s.Credential}
}
return out
}

757
transport/call/vk/joiner.go Normal file
View File

@@ -0,0 +1,757 @@
package vk
import (
"context"
"crypto/tls"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/url"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/common/dialer"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing-box/transport/call/tunnel"
"github.com/sagernet/sing-box/transport/call/wtsignal"
"github.com/sagernet/sing/common/logger"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
)
const (
vkReconnectInitialDelay = time.Second
vkReconnectMaxDelay = 16 * time.Second
vkMaxReconnectAttempts = 10
)
const vkTopologyDirect = "DIRECT"
type vkAuthRottenError struct {
Code string
Msg string
}
func (e *vkAuthRottenError) Error() string {
if e.Msg != "" {
return fmt.Sprintf("auth rotten: %s %s", e.Code, e.Msg)
}
return fmt.Sprintf("auth rotten: %s", e.Code)
}
type VKAuthParams struct {
SessionKey string `json:"sessionKey"`
ApplicationKey string `json:"applicationKey"`
APIBaseURL string `json:"apiBaseURL"`
JoinLink string `json:"joinLink"`
AnonymToken string `json:"anonymToken"`
AppVersion string `json:"appVersion"`
ProtocolVersion string `json:"protocolVersion"`
TunnelMode string `json:"tunnelMode"`
VP8FPS int `json:"vp8Fps"`
VP8Batch int `json:"vp8Batch"`
DualTrack bool `json:"dualTrack"`
}
type VKJoinResponse struct {
Endpoint string `json:"endpoint"`
WtEndpoint string `json:"wt_endpoint"`
Token string `json:"token"`
TurnServer struct {
URLs []string `json:"urls"`
Username string `json:"username"`
Credential string `json:"credential"`
} `json:"turn_server"`
StunServer struct {
URLs []string `json:"urls"`
} `json:"stun_server"`
}
type VKJoiner struct {
logger logger.ContextLogger
OnConnected func(tunnel.DataTunnel)
OnRemoteCandidate func(target int, candidateOrSDP string)
PCConfig common.PeerConnectionConfigurer
AddTracks common.AddTunnelTracksFunc
ReadTrackFn common.ReadTrackFunc
Dialer N.Dialer
DNSRouter adapter.DNSRouter
authParams *VKAuthParams
joinResp *VKJoinResponse
sfu *wtsignal.Conn
vkMu sync.Mutex
vkSeq int
remotePeerID *int64
pc *webrtc.PeerConnection
sampleTrack *webrtc.TrackLocalStaticSample
dc *webrtc.DataChannel
vp8tunnel *tunnel.VP8DataTunnel
sym *tunnel.SymmetricScreenTunnel
producerScreen screenUplink
obf *tunnel.TunnelObfuscator
vp8FPS int
vp8Batch int
dualTrack bool
remoteSet bool
pendingICE []webrtc.ICECandidateInit
configAck tunnel.ConfigAckTracker
reconnectAttempt atomic.Int32
stopCh chan struct{}
stopOnce sync.Once
}
func NewVKJoiner(logger logger.ContextLogger, pcConfig common.PeerConnectionConfigurer, addTracks common.AddTunnelTracksFunc, readTrackFn common.ReadTrackFunc, dialer N.Dialer, dnsRouter adapter.DNSRouter) *VKJoiner {
return &VKJoiner{
logger: logger,
PCConfig: pcConfig,
AddTracks: addTracks,
ReadTrackFn: readTrackFn,
Dialer: dialer,
DNSRouter: dnsRouter,
stopCh: make(chan struct{}),
}
}
func (h *VKJoiner) RunWithParams(jsonParams string) {
var params VKAuthParams
if err := json.Unmarshal([]byte(jsonParams), &params); err != nil {
h.logger.Error(fmt.Sprintf("vk-joiner: failed to parse auth params: %v", err))
return
}
h.authParams = &params
obf, err := tunnel.NewTunnelObfuscator(tunnel.DeriveSecretFromJoinLink(params.JoinLink))
if err != nil {
h.logger.Error(fmt.Sprintf("vk-joiner: obfuscator init failed: %v", err))
return
}
h.obf = obf
h.vp8FPS = params.VP8FPS
h.vp8Batch = params.VP8Batch
// h.dualTrack = params.DualTrack // temporarily disabled for VK joiners
h.logger.Debug("vk-joiner: auth params received")
h.logger.Debug(fmt.Sprintf("vk-joiner: obf key-source=%q localEpoch=0x%08x", params.JoinLink, obf.LocalEpoch()))
h.logger.Debug(fmt.Sprintf("vk-joiner: appVersion=%s protocolVersion=%s vp8Fps=%d vp8Batch=%d",
params.AppVersion, params.ProtocolVersion, params.VP8FPS, params.VP8Batch))
h.logger.Info("vk-joiner: connecting")
if err := h.runOnce(); err != nil {
h.logger.Error(fmt.Sprintf("vk-joiner: %v", err))
return
}
for {
if h.isClosed() {
return
}
h.logger.Info("vk-joiner: tunnel lost")
if !h.waitBeforeRetry(int(h.reconnectAttempt.Load())) {
return
}
attempt := h.reconnectAttempt.Add(1)
if h.isClosed() {
return
}
if int(attempt) > vkMaxReconnectAttempts {
h.logger.Warn(fmt.Sprintf("vk-joiner: gave up after %d consecutive reconnect attempts", vkMaxReconnectAttempts))
return
}
h.logger.Info(fmt.Sprintf("vk-joiner: reconnect attempt #%d", attempt))
if err := h.runOnce(); err != nil {
var authRotten *vkAuthRottenError
if errors.As(err, &authRotten) {
h.logger.Error(fmt.Sprintf("vk-joiner: %v, surrendering", err))
return
}
h.logger.Warn(fmt.Sprintf("vk-joiner: %v, will retry", err))
}
}
}
func (h *VKJoiner) Close() {
h.stopOnce.Do(func() { close(h.stopCh) })
StopCaptchaProxy()
h.vkMu.Lock()
sfu := h.sfu
h.sfu = nil
h.vkMu.Unlock()
if sfu != nil {
sfu.Close()
}
if h.vp8tunnel != nil {
h.vp8tunnel.Stop()
}
if h.pc != nil {
h.pc.Close()
}
}
func (h *VKJoiner) closeTransport() {
h.vkMu.Lock()
sfu := h.sfu
h.vkMu.Unlock()
if sfu != nil {
sfu.Close()
}
}
func (h *VKJoiner) runOnce() error {
h.resetSessionState()
if err := h.joinCall(); err != nil {
return err
}
h.connectSFU()
return nil
}
func (h *VKJoiner) MarkConfigAcked() { h.configAck.Mark() }
func (h *VKJoiner) waitBeforeRetry(attempt int) bool {
delay := common.BackoffWithJitter(attempt, vkReconnectInitialDelay, vkReconnectMaxDelay)
h.logger.Debug(fmt.Sprintf("vk-joiner: waiting %s before reconnect", delay))
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-timer.C:
return !h.isClosed()
case <-h.stopCh:
return false
}
}
func (h *VKJoiner) isClosed() bool {
select {
case <-h.stopCh:
return true
default:
return false
}
}
func (h *VKJoiner) resetSessionState() {
h.vkMu.Lock()
sfu := h.sfu
h.sfu = nil
h.vkSeq = 0
h.vkMu.Unlock()
if sfu != nil {
sfu.Close()
}
if h.sym != nil {
h.sym.Stop()
h.sym = nil
}
if h.vp8tunnel != nil {
h.vp8tunnel.Stop()
h.vp8tunnel = nil
}
h.producerScreen.reset()
if h.dc != nil {
h.dc.Close()
h.dc = nil
}
if h.pc != nil {
h.pc.Close()
h.pc = nil
}
h.sampleTrack = nil
h.remoteSet = false
h.pendingICE = nil
h.remotePeerID = nil
h.joinResp = nil
}
func (h *VKJoiner) dialContext(ctx context.Context, network, addr string) (net.Conn, error) {
return h.Dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
}
func (h *VKJoiner) resolveHost(host string) (string, error) {
if ip := net.ParseIP(host); ip != nil {
return host, nil
}
rd, hasRD := h.Dialer.(dialer.ResolveDialer)
if h.DNSRouter == nil || !hasRD {
return "", fmt.Errorf("no DNS router available to resolve %s", host)
}
addrs, err := h.DNSRouter.Lookup(context.Background(), host, rd.QueryOptions())
if err != nil {
return "", err
}
if len(addrs) == 0 {
return "", fmt.Errorf("no addresses for %s", host)
}
return addrs[0].String(), nil
}
func (h *VKJoiner) joinCall() error {
apiURL := h.authParams.APIBaseURL
parsed, err := url.Parse(apiURL)
if err != nil {
return fmt.Errorf("bad apiBaseURL: %w", err)
}
screenFlag := "false"
if h.dualTrack {
screenFlag = "true"
}
body := url.Values{
"method": {"vchat.joinConversationByLink"},
"session_key": {h.authParams.SessionKey},
"application_key": {h.authParams.ApplicationKey},
"joinLink": {h.authParams.JoinLink},
"anonymToken": {h.authParams.AnonymToken},
"isVideo": {"true"},
"isAudio": {"false"},
"mediaSettings": {`{"isAudioEnabled":false,"isVideoEnabled":true,"isScreenSharingEnabled":` + screenFlag + `}`},
"format": {"json"},
}
client := &http.Client{
Timeout: 15 * time.Second,
Transport: &http.Transport{
TLSClientConfig: &tls.Config{ServerName: parsed.Hostname()},
DialContext: h.dialContext,
},
}
req, err := http.NewRequest("POST", apiURL, strings.NewReader(body.Encode()))
if err != nil {
return fmt.Errorf("new request: %w", err)
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("User-Agent", common.UserAgent)
h.logger.Debug("vk-joiner: calling joinConversationByLink...")
resp, err := client.Do(req)
if err != nil {
return fmt.Errorf("joinConversationByLink: %w", err)
}
defer resp.Body.Close()
raw, err := io.ReadAll(resp.Body)
if err != nil {
return fmt.Errorf("read join response: %w", err)
}
var joinResp VKJoinResponse
if jsonErr := json.Unmarshal(raw, &joinResp); jsonErr != nil {
return fmt.Errorf("decode join response: %w (body: %s)", jsonErr, truncateBody(raw))
}
if joinResp.Endpoint == "" {
if rotten := detectVKAuthRotten(raw); rotten != nil {
return rotten
}
return fmt.Errorf("empty endpoint in join response: %s", truncateBody(raw))
}
h.joinResp = &joinResp
h.logger.Debug(fmt.Sprintf("vk-joiner: joined, turn=%v", joinResp.TurnServer.URLs))
return nil
}
func (h *VKJoiner) connectSFU() {
endpoint := h.joinResp.WtEndpoint
if endpoint == "" {
h.logger.Error("vk-joiner: no wt_endpoint in join response, cannot connect")
return
}
parsed, err := url.Parse(endpoint)
if err != nil {
h.logger.Error(fmt.Sprintf("vk-joiner: bad endpoint URL: %s", common.MaskError(err)))
return
}
hostname := parsed.Hostname()
resolvedIP, err := h.resolveHost(hostname)
if err != nil {
h.logger.Error(fmt.Sprintf("vk-joiner: DNS resolve failed: %s", common.MaskError(err)))
return
}
h.logger.Debug(fmt.Sprintf("vk-joiner: resolved %s -> %s", common.MaskAddr(hostname), common.MaskAddr(resolvedIP)))
capabilities := "2F7F"
wtURL := endpoint +
"&platform=WEB" +
"&appVersion=" + h.authParams.AppVersion +
"&version=" + h.authParams.ProtocolVersion +
"&device=browser&capabilities=" + capabilities + "&clientType=VK&tgt=join&compression=deflate-raw"
sfu, err := wtsignal.Dial(wtURL, hostname, resolvedIP)
if err != nil {
h.logger.Error(fmt.Sprintf("vk-joiner: WebTransport connect failed: %s", common.MaskError(err)))
return
}
h.vkMu.Lock()
h.sfu = sfu
h.vkSeq = 0
h.vkMu.Unlock()
h.logger.Debug("vk-joiner: WebTransport connected")
h.vkSend("update-media-modifiers", map[string]interface{}{
"mediaModifiers": map[string]interface{}{"denoise": true, "denoiseAnn": true},
})
h.vkSend("change-media-settings", map[string]interface{}{
"mediaSettings": map[string]interface{}{
"isAudioEnabled": false, "isVideoEnabled": true,
"isScreenSharingEnabled": h.dualTrack, "isFastScreenSharingEnabled": false,
"isAudioSharingEnabled": false, "isAnimojiEnabled": false,
},
})
h.readLoop()
}
func (h *VKJoiner) vkSend(command string, extra map[string]interface{}) {
h.vkMu.Lock()
defer h.vkMu.Unlock()
if h.sfu == nil {
return
}
h.vkSeq++
extra["command"] = command
extra["sequence"] = h.vkSeq
out, _ := json.Marshal(extra)
h.sfu.Send(out)
h.logger.Debug(fmt.Sprintf("vk-joiner: -> %s", command))
}
func (h *VKJoiner) vkSendTransmitData(participantId int64, payload map[string]interface{}) {
h.vkMu.Lock()
defer h.vkMu.Unlock()
if h.sfu == nil {
return
}
h.vkSeq++
payloadJSON, _ := json.Marshal(payload)
out := fmt.Sprintf(`{"command":"transmit-data","sequence":%d,"participantId":%d,"data":%s}`,
h.vkSeq, participantId, payloadJSON)
h.sfu.Send([]byte(out))
}
func (h *VKJoiner) readLoop() {
h.vkMu.Lock()
sfu := h.sfu
h.vkMu.Unlock()
if sfu == nil {
return
}
for {
msg, err := sfu.Recv()
if err != nil {
h.logger.Debug(fmt.Sprintf("vk-joiner: WebTransport closed: %s", common.MaskError(err)))
return
}
if string(msg) == "ping" {
sfu.Send([]byte("pong"))
continue
}
h.handleVKMessage(msg)
}
}
func (h *VKJoiner) handleVKMessage(raw []byte) {
var msg map[string]interface{}
if err := json.Unmarshal(raw, &msg); err != nil {
return
}
msgType, _ := msg["type"].(string)
switch msgType {
case "notification":
notif, _ := msg["notification"].(string)
switch notif {
case "connection":
h.handleConnection(msg)
case "transmitted-data":
data, _ := msg["data"].(map[string]interface{})
if data != nil {
if pid, ok := msg["participantId"].(float64); ok && h.remotePeerID == nil {
h.onRegisteredPeer(int64(pid))
}
h.onTransmittedData(data)
}
case "registered-peer":
if pid, ok := msg["participantId"].(float64); ok {
h.onRegisteredPeer(int64(pid))
}
case "topology-changed":
topo, _ := msg["topology"].(string)
h.logger.Debug(fmt.Sprintf("vk-joiner: topology: %s", topo))
if topo != "" && topo != vkTopologyDirect {
h.logger.Debug(fmt.Sprintf("vk-joiner: %s topology -> closing transport to reconnect and recover DIRECT", topo))
h.closeTransport()
}
case "participant-joined", "participant-added":
h.logger.Debug(fmt.Sprintf("vk-joiner: <- %s", notif))
case "participant-left":
h.logger.Debug(fmt.Sprintf("vk-joiner: <- %s", notif))
case "hungup":
h.logger.Debug("vk-joiner: peer hungup -> closing transport to reconnect")
h.closeTransport()
}
case "response":
seq, _ := msg["sequence"].(float64)
h.logger.Debug(fmt.Sprintf("vk-joiner: <- response seq=%d", int(seq)))
case "error":
errMsg, _ := msg["message"].(string)
errCode, _ := msg["error"].(string)
h.logger.Warn(fmt.Sprintf("vk-joiner: ERROR: %s %s", errCode, errMsg))
}
}
func (h *VKJoiner) handleConnection(msg map[string]interface{}) {
if conv, ok := msg["conversation"].(map[string]interface{}); ok {
topo, _ := conv["topology"].(string)
h.logger.Debug(fmt.Sprintf("vk-joiner: connection topology=%q", topo))
}
convParams, ok := msg["conversationParams"].(map[string]interface{})
if !ok {
return
}
turn, ok := convParams["turn"].(map[string]interface{})
if !ok {
return
}
urlsRaw, _ := turn["urls"].([]interface{})
var urls []string
for _, u := range urlsRaw {
if s, ok := u.(string); ok {
urls = append(urls, s)
}
}
username, _ := turn["username"].(string)
credential, _ := turn["credential"].(string)
h.joinResp.TurnServer.URLs = urls
h.joinResp.TurnServer.Username = username
h.joinResp.TurnServer.Credential = credential
h.logger.Debug(fmt.Sprintf("vk-joiner: TURN from connection: %v", urls))
if h.pc == nil {
h.initPC()
}
}
func (h *VKJoiner) initPC() {
var iceServers []webrtc.ICEServer
if len(h.joinResp.StunServer.URLs) > 0 {
iceServers = append(iceServers, webrtc.ICEServer{URLs: h.joinResp.StunServer.URLs})
}
if len(h.joinResp.TurnServer.URLs) > 0 {
iceServers = append(iceServers, webrtc.ICEServer{
URLs: h.joinResp.TurnServer.URLs,
Username: h.joinResp.TurnServer.Username,
Credential: h.joinResp.TurnServer.Credential,
})
}
mode := h.authParams.TunnelMode
settingEngine := webrtc.SettingEngine{}
settingEngine.DisableCloseByDTLS(true)
settingEngine.DetachDataChannels()
if h.PCConfig != nil {
h.PCConfig.ConfigureSettingEngine(&settingEngine)
}
pc, err := webrtc.NewAPI(webrtc.WithSettingEngine(settingEngine)).NewPeerConnection(webrtc.Configuration{
ICEServers: iceServers,
})
if err != nil {
h.logger.Error(fmt.Sprintf("vk-joiner: failed to create PC: %v", err))
return
}
h.pc = pc
h.logger.Debug(fmt.Sprintf("vk-joiner: tunnel mode: %s", mode))
if mode == "video" {
h.sampleTrack = h.AddTracks(pc, h.logger, "vk-joiner")
}
negotiated := true
dcID := uint16(2)
dc, err := pc.CreateDataChannel("tunnel", &webrtc.DataChannelInit{
Negotiated: &negotiated,
ID: &dcID,
})
if err != nil {
h.logger.Warn(fmt.Sprintf("vk-joiner: could not create tunnel DC: %v", err))
} else {
h.dc = dc
dc.OnOpen(func() {
h.logger.Debug("vk-joiner: tunnel DC open")
if mode == "dc" {
h.reconnectAttempt.Store(0)
h.logger.Info("vk-joiner: === DC TUNNEL CONNECTED ===")
if h.OnConnected != nil {
h.OnConnected(tunnel.NewDCTunnel(dc, h.obf, common.RTPBufSize, h.logger))
}
}
})
dc.OnClose(func() {
h.logger.Debug("vk-joiner: tunnel DC closed")
})
}
pc.OnICECandidate(func(candidate *webrtc.ICECandidate) {
if candidate == nil {
return
}
h.onLocalICECandidate(candidate)
})
pc.OnConnectionStateChange(func(state webrtc.PeerConnectionState) {
h.logger.Debug(fmt.Sprintf("vk-joiner: PC state: %s", state.String()))
if state == webrtc.PeerConnectionStateFailed || state == webrtc.PeerConnectionStateDisconnected {
h.logger.Debug(fmt.Sprintf("vk-joiner: PC %s, closing transport to trigger reconnect", state.String()))
h.closeTransport()
}
if mode == "video" && state == webrtc.PeerConnectionStateConnected && h.vp8tunnel == nil {
h.reconnectAttempt.Store(0)
h.logger.Info("vk-joiner: === TUNNEL CONNECTED ===")
h.vp8tunnel = tunnel.NewVP8DataTunnel(h.sampleTrack, h.obf, h.logger)
h.vp8tunnel.Start(h.vp8FPS, h.vp8Batch)
var downlink tunnel.DataTunnel = h.vp8tunnel
trackCount := 1
if h.dualTrack {
writer := tunnel.NewScreenWriter(h.obf, "screen-up", h.logger)
writer.Reconfigure(h.vp8tunnel.FPS(), h.vp8tunnel.Batch())
writer.SetSend(h.producerScreen.send)
h.sym = tunnel.NewSymmetricScreenTunnel(h.vp8tunnel, writer, h.obf, h.producerScreen.ready, h.logger)
h.sym.SetTrackCount(2)
downlink = h.sym
trackCount = 2
h.logger.Info("vk-joiner: === SYMMETRIC DUAL-TRACK: camera VP8 + screen DCs ===")
}
vp8tun := h.vp8tunnel
if !h.configAck.Acknowledged() {
acked, cancel := h.configAck.Arm()
go tunnel.SendVP8ConfigUntilAcked(acked, cancel, h.stopCh, vp8tun,
vp8tun.FPS(), vp8tun.Batch(), trackCount, h.logger, "vk-joiner")
h.logger.Debug(fmt.Sprintf("vk-joiner: pushed vp8 config to creator fps=%d batch=%d trackCount=%d", vp8tun.FPS(), vp8tun.Batch(), trackCount))
}
if h.OnConnected != nil {
h.OnConnected(downlink)
}
}
})
if mode == "video" {
pc.OnDataChannel(func(dc *webrtc.DataChannel) {
h.logger.Debug(fmt.Sprintf("vk-joiner: remote DataChannel: label=%q id=%v", dc.Label(), dc.ID()))
if !h.dualTrack {
return
}
switch dc.Label() {
case "consumerScreenShare":
readScreenDataChannel(dc, func(frame []byte) {
if h.sym != nil {
h.sym.HandleScreenFrame(frame)
}
}, h.logger)
case "producerScreenShare":
attachScreenWriterDC(dc, h.producerScreen.attach, h.logger)
}
})
pc.OnTrack(func(track *webrtc.TrackRemote, receiver *webrtc.RTPReceiver) {
h.logger.Debug(fmt.Sprintf("vk-joiner: remote track: codec=%s ssrc=%d", track.Codec().MimeType, track.SSRC()))
go h.ReadTrackFn(track, func(frame []byte) {
if h.vp8tunnel != nil {
h.vp8tunnel.HandleFrame(frame)
}
}, h.logger, "vk-joiner")
})
}
h.logger.Debug("vk-joiner: PC ready, waiting for remote offer")
}
func (h *VKJoiner) onRegisteredPeer(pid int64) {
h.remotePeerID = &pid
h.logger.Debug(fmt.Sprintf("vk-joiner: peer registered: %d", pid))
}
func (h *VKJoiner) onLocalICECandidate(candidate *webrtc.ICECandidate) {
if h.remotePeerID == nil {
return
}
candidateJSON := candidate.ToJSON()
raw, _ := json.Marshal(candidateJSON)
var parsed interface{}
json.Unmarshal(raw, &parsed)
h.vkSendTransmitData(*h.remotePeerID, map[string]interface{}{"candidate": parsed})
}
func (h *VKJoiner) onTransmittedData(data map[string]interface{}) {
if h.pc == nil {
return
}
if candidate, ok := data["candidate"]; ok {
candidateJSON, _ := json.Marshal(candidate)
var candidateInit webrtc.ICECandidateInit
json.Unmarshal(candidateJSON, &candidateInit)
if h.OnRemoteCandidate != nil {
h.OnRemoteCandidate(0, candidateInit.Candidate)
}
if h.remoteSet {
h.pc.AddICECandidate(candidateInit)
} else {
h.pendingICE = append(h.pendingICE, candidateInit)
}
}
if sdp, ok := data["sdp"].(map[string]interface{}); ok {
sdpType, _ := sdp["type"].(string)
sdpStr, _ := sdp["sdp"].(string)
if h.OnRemoteCandidate != nil {
h.OnRemoteCandidate(-1, sdpStr)
}
h.logger.Debug(fmt.Sprintf("vk-joiner: remote SDP: %s", sdpType))
if sdpType == "answer" {
h.pc.SetRemoteDescription(webrtc.SessionDescription{Type: webrtc.SDPTypeAnswer, SDP: sdpStr})
h.remoteSet = true
for _, candidate := range h.pendingICE {
h.pc.AddICECandidate(candidate)
}
h.pendingICE = nil
} else if sdpType == "offer" {
h.pc.SetRemoteDescription(webrtc.SessionDescription{Type: webrtc.SDPTypeOffer, SDP: sdpStr})
h.remoteSet = true
for _, candidate := range h.pendingICE {
h.pc.AddICECandidate(candidate)
}
h.pendingICE = nil
answer, err := h.pc.CreateAnswer(nil)
if err != nil || h.remotePeerID == nil {
h.logger.Warn(fmt.Sprintf("vk-joiner: create answer failed: %v", err))
return
}
h.pc.SetLocalDescription(answer)
sdpJSON, _ := json.Marshal(answer.SDP)
h.vkMu.Lock()
if h.sfu != nil {
h.vkSeq++
raw := fmt.Sprintf(`{"command":"transmit-data","sequence":%d,"participantId":%d,"data":{"sdp":{"sdp":%s,"type":%q},"animojiVersion":2},"participantType":"USER"}`,
h.vkSeq, *h.remotePeerID, sdpJSON, answer.Type.String())
h.sfu.Send([]byte(raw))
h.logger.Debug(fmt.Sprintf("vk-joiner: -> answer (seq=%d)", h.vkSeq))
}
h.vkMu.Unlock()
}
}
}
func detectVKAuthRotten(raw []byte) *vkAuthRottenError {
var generic map[string]interface{}
if err := json.Unmarshal(raw, &generic); err != nil {
return nil
}
errCode, _ := generic["error_code"].(string)
if errCode == "" {
if codeNum, ok := generic["error_code"].(float64); ok {
errCode = fmt.Sprintf("%.0f", codeNum)
}
}
if errCode == "" {
errCode, _ = generic["errorCode"].(string)
}
if errCode == "" {
return nil
}
switch errCode {
case "SESSION_EXPIRED", "AUTH_LOGIN", "SESSION_NOT_FOUND", "INVALID_SESSION_KEY":
errMsg, _ := generic["error_msg"].(string)
return &vkAuthRottenError{Code: errCode, Msg: errMsg}
}
return nil
}
func truncateBody(raw []byte) string {
const maxLen = 200
if len(raw) > maxLen {
return string(raw[:maxLen]) + "..."
}
return string(raw)
}

194
transport/call/vk/p2p.go Normal file
View File

@@ -0,0 +1,194 @@
package vk
import (
"encoding/json"
"fmt"
"github.com/gorilla/websocket"
"github.com/pion/webrtc/v4"
)
type P2PHandler struct {
bridge *Bridge
remotePeerId *int64
pendingOffer *webrtc.SessionDescription
pendingCandidates []webrtc.ICECandidateInit
connected bool
}
func NewP2PHandler(bridge *Bridge) *P2PHandler {
return &P2PHandler{bridge: bridge}
}
func (p *P2PHandler) Init() {
relay := p.bridge.relay
if err := relay.Init(p.bridge.iceServers); err != nil {
p.bridge.logger.Error(fmt.Sprintf("[p2p] Init relay failed: %v", err))
return
}
p.setupCallbacks()
offer, err := relay.CreateOffer()
if err != nil {
p.bridge.logger.Error(fmt.Sprintf("[p2p] Create offer failed: %v", err))
return
}
p.pendingOffer = &offer
p.bridge.logger.Debug(fmt.Sprintf("[p2p] Offer ready, SDP length: %d", len(offer.SDP)))
}
func (p *P2PHandler) Reset() {
p.bridge.logger.Debug("[p2p] Resetting Pion PC...")
p.connected = false
p.bridge.relay.Close()
p.bridge.relay = p.bridge.newRelay()
relay := p.bridge.relay
if err := relay.Init(p.bridge.iceServers); err != nil {
p.bridge.logger.Warn(fmt.Sprintf("[p2p] Reset init failed: %v", err))
return
}
p.setupCallbacks()
offer, err := relay.CreateOffer()
if err != nil {
p.bridge.logger.Warn(fmt.Sprintf("[p2p] Reset create-offer failed: %v", err))
return
}
p.pendingOffer = &offer
p.pendingCandidates = nil
p.bridge.logger.Debug(fmt.Sprintf("[p2p] New offer ready after reset, SDP length: %d", len(offer.SDP)))
}
func (p *P2PHandler) OnRegisteredPeer(participantId int64) {
oldPeer := p.remotePeerId
p.remotePeerId = &participantId
p.bridge.logger.Debug(fmt.Sprintf("[p2p] Peer registered: %d (connected=%v)", participantId, p.connected))
if oldPeer != nil && (p.pendingOffer == nil) {
if *oldPeer != participantId {
p.bridge.logger.Debug(fmt.Sprintf("[p2p] New peer %d replacing old peer %d, resetting", participantId, *oldPeer))
} else {
p.bridge.logger.Debug(fmt.Sprintf("[p2p] Same peer %d re-registered, no pending offer, resetting", participantId))
}
p.Reset()
}
p.sendOfferToPeer(participantId)
}
func (p *P2PHandler) OnTransmittedData(data map[string]interface{}) {
if cand, ok := data["candidate"]; ok {
p.bridge.logger.Debug("[p2p] Remote ICE candidate")
candJSON, _ := json.Marshal(cand)
var candInit webrtc.ICECandidateInit
json.Unmarshal(candJSON, &candInit)
p.bridge.relay.AddICECandidate(candInit)
}
if sdp, ok := data["sdp"].(map[string]interface{}); ok {
sdpType, _ := sdp["type"].(string)
sdpStr, _ := sdp["sdp"].(string)
p.bridge.logger.Debug(fmt.Sprintf("[p2p] Remote SDP: %s", sdpType))
if sdpType == "answer" {
p.bridge.relay.SetRemoteDescription(webrtc.SDPTypeAnswer, sdpStr)
} else if sdpType == "offer" {
p.bridge.relay.SetRemoteDescription(webrtc.SDPTypeOffer, sdpStr)
answer, err := p.bridge.relay.CreateAnswer()
if err == nil && p.remotePeerId != nil {
p.bridge.vkSend("transmit-data", map[string]interface{}{
"participantId": *p.remotePeerId,
"data": map[string]interface{}{"sdp": map[string]interface{}{
"type": answer.Type.String(), "sdp": answer.SDP,
}},
})
}
}
}
}
func (p *P2PHandler) OnPionICECandidate(data json.RawMessage) {
if p.remotePeerId != nil {
var cand interface{}
json.Unmarshal(data, &cand)
p.bridge.vkSend("transmit-data", map[string]interface{}{
"participantId": *p.remotePeerId,
"data": map[string]interface{}{"candidate": cand},
})
} else {
var candInit webrtc.ICECandidateInit
json.Unmarshal(data, &candInit)
p.pendingCandidates = append(p.pendingCandidates, candInit)
}
}
func (p *P2PHandler) OnConnectionState(state string) {
switch state {
case "connected":
p.connected = true
p.bridge.logger.Info("TUNNEL CONNECTED")
case "disconnected":
p.connected = false
p.bridge.logger.Debug("[p2p] Connection disconnected, kicking peer")
p.kickRemotePeer()
case "failed":
p.connected = false
p.bridge.logger.Debug("[p2p] Connection failed, removing stale peer")
p.kickRemotePeer()
case "closed":
p.connected = false
p.bridge.logger.Debug("[p2p] Connection closed, kicking peer")
p.kickRemotePeer()
}
}
func (p *P2PHandler) setupCallbacks() {
p.bridge.relay.OnICECandidate(func(cand *webrtc.ICECandidate) {
if cand == nil {
return
}
p.bridge.logger.Debug(fmt.Sprintf("[p2p] ICE candidate: type=%s proto=%s", cand.Typ.String(), cand.Protocol.String()))
candJSON := cand.ToJSON()
raw, _ := json.Marshal(candJSON)
p.OnPionICECandidate(raw)
})
p.bridge.relay.OnConnectionStateChange(func(state webrtc.PeerConnectionState) {
p.OnConnectionState(state.String())
})
}
func (p *P2PHandler) kickRemotePeer() {
if p.remotePeerId != nil {
p.bridge.vkSend("remove-participant", map[string]interface{}{
"participantId": *p.remotePeerId,
"ban": false,
})
}
}
func (p *P2PHandler) sendOfferToPeer(participantId int64) {
offer := p.pendingOffer
candidates := p.pendingCandidates
p.pendingOffer = nil
p.pendingCandidates = nil
if offer != nil {
p.bridge.logger.Debug(fmt.Sprintf("[p2p] Sending offer to peer %d", participantId))
sdpStr, _ := json.Marshal(offer.SDP)
p.bridge.mu.Lock()
p.bridge.vkSeq++
seq := p.bridge.vkSeq
raw := fmt.Sprintf(`{"command":"transmit-data","sequence":%d,"participantId":%d,"data":{"sdp":{"type":%q,"sdp":%s}}}`,
seq, participantId, offer.Type.String(), sdpStr)
if p.bridge.vkWs != nil {
p.bridge.vkWs.WriteMessage(websocket.TextMessage, []byte(raw))
}
p.bridge.mu.Unlock()
p.bridge.logger.Debug("[vk-ws] -> transmit-data (offer)")
}
for _, cand := range candidates {
candJSON, _ := json.Marshal(cand)
var c interface{}
json.Unmarshal(candJSON, &c)
p.bridge.vkSend("transmit-data", map[string]interface{}{
"participantId": participantId,
"data": map[string]interface{}{"candidate": c},
})
}
if len(candidates) > 0 {
p.bridge.logger.Debug(fmt.Sprintf("[p2p] Flushed %d ICE candidates", len(candidates)))
}
}

468
transport/call/vk/relay.go Normal file
View File

@@ -0,0 +1,468 @@
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
}
}

View File

@@ -0,0 +1,101 @@
package vk
import (
"errors"
"io"
"sync"
"github.com/pion/datachannel"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing/common/logger"
)
var errScreenNotReady = errors.New("screen DC not ready")
type screenUplink struct {
mu sync.Mutex
raw io.WriteCloser
}
func (u *screenUplink) attach(raw io.WriteCloser) {
u.mu.Lock()
u.raw = raw
u.mu.Unlock()
}
func (u *screenUplink) ready() bool {
u.mu.Lock()
defer u.mu.Unlock()
return u.raw != nil
}
func (u *screenUplink) send(b []byte) error {
u.mu.Lock()
raw := u.raw
u.mu.Unlock()
if raw == nil {
return errScreenNotReady
}
_, err := raw.Write(b)
if err != nil {
u.mu.Lock()
if u.raw == raw {
u.raw = nil
}
u.mu.Unlock()
}
return err
}
func (u *screenUplink) reset() {
u.mu.Lock()
if u.raw != nil {
u.raw.Close()
u.raw = nil
}
u.mu.Unlock()
}
func readScreenDataChannel(dc *webrtc.DataChannel, handler func([]byte), logger logger.ContextLogger) {
dc.OnOpen(func() {
var raw datachannel.ReadWriteCloser
raw, err := dc.Detach()
if err != nil {
logger.Warn("[vk-joiner] screen DC detach failed, using OnMessage")
dc.OnMessage(func(m webrtc.DataChannelMessage) {
if !m.IsString && len(m.Data) > 0 {
frame := make([]byte, len(m.Data))
copy(frame, m.Data)
handler(frame)
}
})
return
}
logger.Debug("[vk-joiner] screen DC attached for reading")
buf := make([]byte, 65536)
for {
n, isString, rerr := raw.ReadDataChannel(buf)
if rerr != nil {
return
}
if isString || n == 0 {
continue
}
frame := make([]byte, n)
copy(frame, buf[:n])
handler(frame)
}
})
}
func attachScreenWriterDC(dc *webrtc.DataChannel, onRaw func(io.WriteCloser), logger logger.ContextLogger) {
dc.OnOpen(func() {
raw, err := dc.Detach()
if err != nil {
logger.Warn("[vk-joiner] screen writer DC detach failed")
return
}
logger.Debug("[vk-joiner] screen DC attached for writing")
onRaw(raw)
})
}

View File

@@ -0,0 +1,244 @@
package vk
import (
"encoding/json"
"fmt"
"io"
"math/rand"
"net/http"
"net/url"
"strings"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
type vkAuthConfig struct {
AppID string `json:"appID"`
ApiVersion string `json:"apiVersion"`
AppVersion string `json:"appVersion"`
ProtocolVersion string `json:"protocolVersion"`
PublicKey string `json:"publicKey,omitempty"`
OkJoinLink string `json:"okJoinLink,omitempty"`
}
type vkCaptchaError struct {
captchaSid string
redirectURI string
captchaTs string
captchaAttempt string
}
func RunVKAuth(dialer N.Dialer, joinLink, displayName string, logger logger.ContextLogger) (string, error) {
client := common.HttpClient(dialer)
httpPost := func(targetURL string, form url.Values, extraHeaders map[string]string) (map[string]interface{}, error) {
req, _ := http.NewRequest("POST", targetURL, strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("User-Agent", common.UserAgent)
req.Header.Set("Origin", "https://vk.ru")
req.Header.Set("Referer", "https://vk.ru/")
for k, v := range extraHeaders {
req.Header.Set(k, v)
}
resp, err := client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return nil, err
}
var result map[string]interface{}
if err := json.Unmarshal(body, &result); err != nil {
return nil, fmt.Errorf("json: %w (body: %s)", err, string(body[:minInt(len(body), 200)]))
}
return result, nil
}
cfg := &vkAuthConfig{
AppID: "6287487",
ApiVersion: "5.282",
AppVersion: "1.1",
ProtocolVersion: "5",
}
logger.Debug(fmt.Sprintf("vk-auth: appID=%s api=%s appVersion=%s proto=%s", cfg.AppID, cfg.ApiVersion, cfg.AppVersion, cfg.ProtocolVersion))
logger.Info("vk-auth: getting anonymous token")
anonResp, err := httpPost("https://login.vk.ru/?act=get_anonym_token", url.Values{
"client_id": {cfg.AppID},
}, nil)
if err != nil {
return "", fmt.Errorf("get_anonym_token: %w", err)
}
dataMap, _ := anonResp["data"].(map[string]interface{})
accessToken, _ := dataMap["access_token"].(string)
if accessToken == "" {
return "", fmt.Errorf("empty access_token: %v", anonResp)
}
logger.Debug("vk-auth: anon token OK")
auth := map[string]string{"Authorization": "Bearer " + accessToken}
logger.Info("vk-auth: getting call settings")
settingsResp, err := httpPost("https://api.vk.ru/method/calls.getSettings", url.Values{
"v": {cfg.ApiVersion},
}, auth)
if err != nil {
return "", fmt.Errorf("calls.getSettings: %w", err)
}
if respObj, ok := settingsResp["response"].(map[string]interface{}); ok {
if settings, ok := respObj["settings"].(map[string]interface{}); ok {
if pk, ok := settings["public_key"].(string); ok {
cfg.PublicKey = pk
}
}
}
logger.Debug(fmt.Sprintf("vk-auth: publicKey=%s", cfg.PublicKey))
logger.Info("vk-auth: getting call preview")
previewResp, err := httpPost("https://api.vk.ru/method/calls.getCallPreview", url.Values{
"v": {cfg.ApiVersion},
"vk_join_link": {joinLink},
}, auth)
if err == nil {
if respObj, ok := previewResp["response"].(map[string]interface{}); ok {
if okLink, ok := respObj["ok_join_link"].(string); ok {
cfg.OkJoinLink = okLink
}
}
}
logger.Info("vk-auth: getting call token")
callParams := url.Values{
"v": {cfg.ApiVersion},
"vk_join_link": {joinLink},
"name": {displayName},
}
var callToken string
var apiBaseURL string
var okJoinLink string
for attempt := 0; attempt < 5; attempt++ {
callResp, err := httpPost("https://api.vk.ru/method/calls.getAnonymousToken", callParams, auth)
if err != nil {
return "", fmt.Errorf("getAnonymousToken: %w", err)
}
if errObj, hasErr := callResp["error"].(map[string]interface{}); hasErr {
errCode, _ := errObj["error_code"].(float64)
if int(errCode) == 14 {
captchaErr := parseVKCaptchaError(errObj)
if captchaErr == nil {
return "", fmt.Errorf("captcha error missing fields: %v", errObj)
}
logger.Info("vk-auth: captcha required")
proxyPort := StartCaptchaProxy(captchaErr.redirectURI, dialer)
if proxyPort == 0 {
return "", fmt.Errorf("failed to start captcha proxy")
}
logger.Notice(fmt.Sprintf("vk-auth: solve the captcha to continue: http://127.0.0.1:%d/", proxyPort))
successToken := GetCaptchaResult()
StopCaptchaProxy()
if successToken == "" {
return "", fmt.Errorf("captcha timed out")
}
logger.Info("vk-auth: captcha solved, retrying")
captchaAttempt := captchaErr.captchaAttempt
if captchaAttempt == "" || captchaAttempt == "0" {
captchaAttempt = "1"
}
callParams = url.Values{
"v": {cfg.ApiVersion},
"vk_join_link": {joinLink},
"name": {displayName},
"captcha_key": {""},
"captcha_sid": {captchaErr.captchaSid},
"is_sound_captcha": {"0"},
"success_token": {successToken},
"captcha_ts": {captchaErr.captchaTs},
"captcha_attempt": {captchaAttempt},
}
continue
}
return "", fmt.Errorf("VK API error: %v", errObj)
}
respMap, ok := callResp["response"].(map[string]interface{})
if !ok {
return "", fmt.Errorf("unexpected response: %v", callResp)
}
callToken, _ = respMap["token"].(string)
apiBaseURL, _ = respMap["api_base_url"].(string)
okJoinLink, _ = respMap["ok_join_link"].(string)
break
}
if callToken == "" {
return "", fmt.Errorf("failed to get call token")
}
logger.Info("vk-auth: authenticating with OK.ru")
baseURL := strings.TrimRight(apiBaseURL, "/")
if !strings.HasSuffix(baseURL, "/fb.do") {
baseURL += "/fb.do"
}
deviceID := fmt.Sprintf("%d", rand.Int63n(9e18))
sessionData, _ := json.Marshal(map[string]interface{}{
"version": 2,
"device_id": deviceID,
"client_version": cfg.AppVersion,
"client_type": "SDK_JS",
})
okResp, err := httpPost(baseURL, url.Values{
"method": {"auth.anonymLogin"},
"session_data": {string(sessionData)},
"application_key": {cfg.PublicKey},
"format": {"json"},
}, nil)
if err != nil {
return "", fmt.Errorf("anonymLogin: %w", err)
}
sessionKey, _ := okResp["session_key"].(string)
if sessionKey == "" {
return "", fmt.Errorf("missing session_key: %v", okResp)
}
logger.Debug("vk-auth: OK.ru session OK")
finalJoinLink := okJoinLink
if finalJoinLink == "" {
finalJoinLink = cfg.OkJoinLink
}
if finalJoinLink == "" {
finalJoinLink = joinLink
}
result := map[string]string{
"sessionKey": sessionKey,
"applicationKey": cfg.PublicKey,
"apiBaseURL": baseURL,
"joinLink": finalJoinLink,
"anonymToken": callToken,
"appVersion": cfg.AppVersion,
"protocolVersion": cfg.ProtocolVersion,
}
jsonBytes, _ := json.Marshal(result)
logger.Debug("vk-auth: done")
return string(jsonBytes), nil
}
func parseVKCaptchaError(errObj map[string]interface{}) *vkCaptchaError {
redirectURI, _ := errObj["redirect_uri"].(string)
if redirectURI == "" {
return nil
}
captchaSid := ""
if sid, ok := errObj["captcha_sid"].(string); ok {
captchaSid = sid
} else if sidNum, ok := errObj["captcha_sid"].(float64); ok {
captchaSid = fmt.Sprintf("%.0f", sidNum)
}
captchaTs, _ := errObj["captcha_ts"].(string)
captchaAttempt, _ := errObj["captcha_attempt"].(string)
return &vkCaptchaError{
captchaSid: captchaSid,
redirectURI: redirectURI,
captchaTs: captchaTs,
captchaAttempt: captchaAttempt,
}
}
func minInt(a, b int) int {
if a < b {
return a
}
return b
}

View File

@@ -0,0 +1,363 @@
package wbstream
import (
"bytes"
"crypto/rand"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"github.com/sagernet/sing-box/transport/call/common"
)
const (
APIBase = "https://stream.wb.ru"
Origin = "https://stream.wb.ru"
)
var WBStreamCookieAllowlist = []string{
"wbx-refresh",
"x_wbaas_token",
"_wbauid",
"wbx-validation-key",
}
var ModeratorPermissions = []string{
"ROOM_PERMISSION_SEND_CHAT",
"ROOM_PERMISSION_SHARE_AUDIO",
"ROOM_PERMISSION_SHARE_SCREEN",
"ROOM_PERMISSION_SHARE_VIDEO",
"ROOM_PERMISSION_MODIFY_PERMISSIONS",
"ROOM_PERMISSION_MODERATE_ROOM",
"ROOM_PERMISSION_CALL_DATA_ACCESS",
"ROOM_PERMISSION_LOCAL_RECORD",
}
type guestRegisterRequest struct {
DisplayName string `json:"displayName"`
Device guestDeviceCfg `json:"device"`
}
type guestDeviceCfg struct {
DeviceName string `json:"deviceName"`
DeviceType string `json:"deviceType"`
}
type guestRegisterResponse struct {
AccessToken string `json:"accessToken"`
}
type createRoomRequest struct {
RoomType string `json:"roomType"`
RoomPrivacy string `json:"roomPrivacy"`
}
type createRoomResponse struct {
RoomID string `json:"roomId"`
}
type connectionDetailsResponse struct {
RoomToken string `json:"roomToken"`
ServerURL string `json:"serverUrl"`
}
type cookieTransport struct {
base http.RoundTripper
cookie string
}
type slideV3Response struct {
Payload struct {
AccessToken string `json:"access_token"`
} `json:"payload"`
}
func ParseRoomID(input string) string {
trimmed := strings.TrimSpace(input)
if trimmed == "" {
return ""
}
if rest, ok := strings.CutPrefix(trimmed, "wbstream://"); ok {
return strings.Trim(rest, "/")
}
if strings.HasPrefix(trimmed, "http://") || strings.HasPrefix(trimmed, "https://") {
u, err := url.Parse(trimmed)
if err == nil {
parts := strings.Split(strings.Trim(u.Path, "/"), "/")
for i := 0; i < len(parts)-1; i++ {
if parts[i] == "room" && parts[i+1] != "" {
return parts[i+1]
}
}
}
}
return strings.Trim(trimmed, "/")
}
func RegisterGuest(client *http.Client, displayName string) (string, error) {
body, _ := json.Marshal(guestRegisterRequest{
DisplayName: displayName,
Device: guestDeviceCfg{
DeviceName: "Linux",
DeviceType: "PARTICIPANT_DEVICE_TYPE_WEB_DESKTOP",
},
})
req, err := http.NewRequest(http.MethodPost, APIBase+"/auth/api/v1/auth/user/guest-register", bytes.NewReader(body))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
resp, err := httpDo(client, req)
if err != nil {
return "", err
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("guest-register: status %d: %s", resp.StatusCode, string(raw))
}
var r guestRegisterResponse
if err := json.Unmarshal(raw, &r); err != nil {
return "", fmt.Errorf("guest-register decode: %w", err)
}
return r.AccessToken, nil
}
func CreateRoom(client *http.Client, accessToken string) (string, error) {
body, _ := json.Marshal(createRoomRequest{
RoomType: "ROOM_TYPE_ALL_ON_SCREEN",
RoomPrivacy: "ROOM_PRIVACY_FREE",
})
req, err := http.NewRequest(http.MethodPost, APIBase+"/api-room/api/v2/room", bytes.NewReader(body))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/json")
setBearer(req, accessToken)
resp, err := httpDo(client, req)
if err != nil {
return "", err
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated {
return "", fmt.Errorf("create-room: status %d: %s", resp.StatusCode, string(raw))
}
var r createRoomResponse
if err := json.Unmarshal(raw, &r); err != nil {
return "", fmt.Errorf("create-room decode: %w", err)
}
return r.RoomID, nil
}
func JoinRoom(client *http.Client, accessToken, roomID string) error {
url := fmt.Sprintf("%s/api-room/api/v1/room/%s/join", APIBase, roomID)
req, err := http.NewRequest(http.MethodPost, url, bytes.NewReader([]byte("{}")))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
setBearer(req, accessToken)
resp, err := httpDo(client, req)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
raw, _ := io.ReadAll(resp.Body)
return fmt.Errorf("join-room: status %d: %s", resp.StatusCode, string(raw))
}
return nil
}
func GetConnectionDetails(client *http.Client, accessToken, roomID, displayName string) (string, string, error) {
detailsURL := fmt.Sprintf("%s/api-room-manager/v2/room/%s/connection-details?deviceType=PARTICIPANT_DEVICE_TYPE_WEB_DESKTOP&displayName=%s",
APIBase, roomID, url.QueryEscape(displayName))
req, err := http.NewRequest(http.MethodGet, detailsURL, nil)
if err != nil {
return "", "", err
}
setBearer(req, accessToken)
resp, err := httpDo(client, req)
if err != nil {
return "", "", err
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return "", "", fmt.Errorf("connection-details: status %d: %s", resp.StatusCode, string(raw))
}
var r connectionDetailsResponse
if err := json.Unmarshal(raw, &r); err != nil {
return "", "", fmt.Errorf("connection-details decode: %w", err)
}
return r.RoomToken, r.ServerURL, nil
}
func AuthAndGetToken(client *http.Client, roomID, displayName string) (string, string, string, string, error) {
accessToken, err := RegisterGuest(client, displayName)
if err != nil {
return "", "", "", "", fmt.Errorf("register guest: %w", err)
}
return joinAndGetDetails(client, accessToken, roomID, displayName)
}
func AuthAsLoggedIn(client *http.Client, cookieHeader, accessToken, roomID, displayName string) (string, string, string, string, error) {
if cookieHeader == "" && accessToken == "" {
return "", "", "", "", fmt.Errorf("cookies or access token required for logged-in auth")
}
client = clientWithCookies(client, cookieHeader)
return joinAndGetDetails(client, accessToken, roomID, displayName)
}
func RefreshAccessToken(client *http.Client, cookieHeader, deviceID string) (string, error) {
req, err := http.NewRequest(http.MethodPost, "https://auth-stream.wb.ru/v2/auth/slide-v3", bytes.NewReader(nil))
if err != nil {
return "", err
}
if deviceID == "" {
deviceID = newRequestID()
}
req.Header.Set("wb-apptype", "web")
req.Header.Set("X-Real-IP", "")
req.Header.Set("deviceId", deviceID)
req.Header.Set("X-Request-ID", newRequestID())
req.Header.Set("Origin", Origin)
req.Header.Set("Referer", Origin+"/")
req.Header.Set("Cookie", cookieHeader)
req.Header.Set("User-Agent", common.UserAgent)
if client == nil {
client = http.DefaultClient
}
resp, err := client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
raw, _ := io.ReadAll(resp.Body)
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("slide-v3: status %d: %s", resp.StatusCode, string(raw))
}
var r slideV3Response
if err := json.Unmarshal(raw, &r); err != nil {
return "", fmt.Errorf("slide-v3 decode: %w", err)
}
if r.Payload.AccessToken == "" {
return "", fmt.Errorf("slide-v3: empty access_token in response: %s", string(raw))
}
return r.Payload.AccessToken, nil
}
func SetParticipantPermissions(client *http.Client, accessToken, roomID, participantID string, permissions []string) error {
setURL := fmt.Sprintf("%s/api-room-manager/api/v1/room/%s/participant/%s/set-permissions", APIBase, roomID, participantID)
body, err := json.Marshal(map[string]any{"permissions": permissions})
if err != nil {
return err
}
req, err := http.NewRequest(http.MethodPut, setURL, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
setBearer(req, accessToken)
resp, err := httpDo(client, req)
if err != nil {
return err
}
defer resp.Body.Close()
respBody, _ := io.ReadAll(resp.Body)
if resp.StatusCode >= 400 {
return fmt.Errorf("set-permissions %s -> %d %s", participantID, resp.StatusCode, string(respBody))
}
return nil
}
func KickParticipant(client *http.Client, accessToken, roomID, participantID string) error {
if client == nil {
client = http.DefaultClient
}
kickURL := fmt.Sprintf("%s/api-room-manager/api/v1/room/%s/participant/%s/kick", APIBase, roomID, participantID)
req, err := http.NewRequest("DELETE", kickURL, strings.NewReader("{}"))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+accessToken)
req.Header.Set("User-Agent", common.UserAgent)
resp, err := client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode >= 400 {
return fmt.Errorf("kick %s -> %d %s", participantID, resp.StatusCode, string(body))
}
return nil
}
func (t *cookieTransport) RoundTrip(req *http.Request) (*http.Response, error) {
req.Header.Set("Cookie", t.cookie)
base := t.base
if base == nil {
base = http.DefaultTransport
}
return base.RoundTrip(req)
}
func httpDo(client *http.Client, req *http.Request) (*http.Response, error) {
req.Header.Set("User-Agent", common.UserAgent)
if client == nil {
client = http.DefaultClient
}
return client.Do(req)
}
func clientWithCookies(client *http.Client, cookieHeader string) *http.Client {
if cookieHeader == "" {
return client
}
if client == nil {
client = &http.Client{}
}
wrapped := *client
wrapped.Transport = &cookieTransport{base: client.Transport, cookie: cookieHeader}
return &wrapped
}
func setBearer(req *http.Request, accessToken string) {
if accessToken != "" {
req.Header.Set("Authorization", "Bearer "+accessToken)
}
}
func newRequestID() string {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "00000000-0000-0000-0000-000000000000"
}
b[6] = (b[6] & 0x0f) | 0x40
b[8] = (b[8] & 0x3f) | 0x80
return fmt.Sprintf("%x-%x-%x-%x-%x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:16])
}
func joinAndGetDetails(client *http.Client, accessToken, roomID, displayName string) (string, string, string, string, error) {
var err error
if roomID == "" {
roomID, err = CreateRoom(client, accessToken)
if err != nil {
return "", "", "", "", fmt.Errorf("create room: %w", err)
}
}
if err := JoinRoom(client, accessToken, roomID); err != nil {
return "", "", "", "", fmt.Errorf("join room: %w", err)
}
roomToken, serverURL, err := GetConnectionDetails(client, accessToken, roomID, displayName)
if err != nil {
return "", "", "", "", fmt.Errorf("get connection details: %w", err)
}
return roomID, roomToken, accessToken, serverURL, nil
}

View File

@@ -0,0 +1,185 @@
package wbstream
import (
"context"
"fmt"
"net/http"
"time"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing-box/transport/call/tunnel"
"github.com/sagernet/sing/common/logger"
N "github.com/sagernet/sing/common/network"
)
func ConnectCreator(ctx context.Context, cookieStr, roomID, mode string, readBuf int, dialer N.Dialer, logger logger.ContextLogger) (*tunnel.RelayBridge, string, error) {
deviceID := common.CookieValue(cookieStr, "__wb_device_id")
if deviceID == "" {
return nil, "", fmt.Errorf("wbstream: cookies missing __wb_device_id")
}
cookieHeader := common.FilterCookies(cookieStr, WBStreamCookieAllowlist)
httpClient := common.HttpClient(dialer)
bearer, err := RefreshAccessToken(httpClient, cookieHeader, deviceID)
if err != nil {
return nil, "", fmt.Errorf("wbstream: slide-v3 refresh: %w", err)
}
requestedRoom := ParseRoomID(roomID)
resolvedRoomID, roomToken, accessToken, serverURL, err := AuthAsLoggedIn(httpClient, cookieHeader, bearer, requestedRoom, "Creator")
if err != nil {
return nil, "", fmt.Errorf("wbstream: auth: %w", err)
}
if readBuf <= 0 {
readBuf = 32768
}
if mode == "" {
mode = TunnelModeDC
}
obf, err := tunnel.NewTunnelObfuscator(tunnel.DeriveSecretFromJoinLink(resolvedRoomID))
if err != nil {
return nil, "", fmt.Errorf("wbstream: obfuscator init: %w", err)
}
joinSession := func(token, access, server string) (*Session, <-chan tunnel.DataTunnel) {
tunCh := make(chan tunnel.DataTunnel, 1)
sess := NewSession(SessionConfig{
RoomToken: token,
ServerURL: server,
DisplayName: "Creator",
TunnelMode: mode,
Obfuscator: obf,
Logger: logger,
Dialer: dialer,
RoomID: resolvedRoomID,
AccessToken: access,
ReadBuf: readBuf,
})
sess.OnConnected = func(tun tunnel.DataTunnel) {
select {
case tunCh <- tun:
default:
}
}
return sess, tunCh
}
sess, tunCh := joinSession(roomToken, accessToken, serverURL)
if err := sess.Start(); err != nil {
return nil, "", fmt.Errorf("wbstream: session start: %w", err)
}
var firstTun tunnel.DataTunnel
select {
case firstTun = <-tunCh:
case <-ctx.Done():
sess.Close()
return nil, "", ctx.Err()
case <-time.After(60 * time.Second):
sess.Close()
return nil, "", fmt.Errorf("wbstream: creator tunnel timed out")
}
relay := tunnel.NewRelayBridge(firstTun, "creator", bridgeReadBufFor(firstTun, readBuf), dialer, logger)
go creatorReconnectLoop(ctx, relay, sess, joinSession, httpClient, cookieHeader, deviceID, resolvedRoomID, readBuf, logger)
return relay, APIBase + "/room/" + resolvedRoomID, nil
}
func ConnectJoiner(ctx context.Context, roomID, displayName, mode string, readBuf int, dialer N.Dialer, dnsRouter adapter.DNSRouter, logger logger.ContextLogger) (tunnel.DataTunnel, error) {
roomID = ParseRoomID(roomID)
if displayName == "" {
displayName = "Joiner"
}
if mode == "" {
mode = TunnelModeDC
}
joiner := NewWBStreamJoiner(logger, dialer, dnsRouter, nil)
tunCh := make(chan tunnel.DataTunnel, 1)
joiner.OnConnected = func(tun tunnel.DataTunnel) {
select {
case tunCh <- tun:
default:
}
}
params := fmt.Sprintf(`{"roomId":%q,"displayName":%q,"tunnelMode":%q}`, roomID, displayName, mode)
go joiner.RunWithParams(params)
select {
case tun := <-tunCh:
return tun, nil
case <-ctx.Done():
joiner.Close()
return nil, ctx.Err()
}
}
func creatorReconnectLoop(
ctx context.Context,
relay *tunnel.RelayBridge,
sess *Session,
joinSession func(token, access, server string) (*Session, <-chan tunnel.DataTunnel),
httpClient *http.Client,
cookieHeader, deviceID, roomID string,
readBuf int,
logger logger.ContextLogger,
) {
current := sess
for {
select {
case <-current.Done():
case <-ctx.Done():
return
}
current.Close()
if relay.IsClosed() {
return
}
logger.Debug("wbstream: creator session ended, rejoining")
var newTun tunnel.DataTunnel
for {
select {
case <-ctx.Done():
return
case <-time.After(3 * time.Second):
}
if relay.IsClosed() {
return
}
bearer, err := RefreshAccessToken(httpClient, cookieHeader, deviceID)
if err != nil {
logger.Warn(fmt.Sprintf("wbstream: rejoin token refresh failed: %v, retrying", err))
continue
}
_, roomToken, accessToken, serverURL, err := AuthAsLoggedIn(httpClient, cookieHeader, bearer, roomID, "Creator")
if err != nil {
logger.Warn(fmt.Sprintf("wbstream: rejoin auth failed: %v, retrying", err))
continue
}
newSess, tunCh := joinSession(roomToken, accessToken, serverURL)
if err := newSess.Start(); err != nil {
logger.Warn(fmt.Sprintf("wbstream: rejoin session start failed: %v, retrying", err))
continue
}
select {
case newTun = <-tunCh:
case <-ctx.Done():
newSess.Close()
return
case <-time.After(60 * time.Second):
logger.Warn("wbstream: rejoin tunnel timed out, retrying")
newSess.Close()
continue
}
current = newSess
break
}
relay.SwapTunnel(newTun)
logger.Info(fmt.Sprintf("wbstream: creator tunnel reconnected, buf=%d", bridgeReadBufFor(newTun, readBuf)))
}
}
func bridgeReadBufFor(tun tunnel.DataTunnel, readBuf int) int {
switch tun.(type) {
case *tunnel.DCTunnel, *tunnel.MultiTrackKCPTunnel:
return readBuf
}
return common.VP8BufSize
}

View File

@@ -0,0 +1,53 @@
package wbstream
import (
"github.com/pion/datachannel"
"github.com/sagernet/sing-box/transport/call/livekit"
)
type dataPacketWrapper struct {
inner datachannel.ReadWriteCloser
kind int
}
func (w *dataPacketWrapper) ReadDataChannel(p []byte) (int, bool, error) {
buf := make([]byte, len(p))
for {
n, isString, err := w.inner.ReadDataChannel(buf)
if err != nil {
return 0, false, err
}
if n == 0 {
continue
}
payload, ok := livekit.DecodeDataPacketUser(buf[:n])
if !ok || len(payload) == 0 {
continue
}
copied := copy(p, payload)
return copied, isString, nil
}
}
func (w *dataPacketWrapper) WriteDataChannel(p []byte, isString bool) (int, error) {
wire := livekit.EncodeDataPacketUser(p, w.kind)
if _, err := w.inner.WriteDataChannel(wire, isString); err != nil {
return 0, err
}
return len(p), nil
}
func (w *dataPacketWrapper) Read(p []byte) (int, error) {
n, _, err := w.ReadDataChannel(p)
return n, err
}
func (w *dataPacketWrapper) Write(p []byte) (int, error) {
return w.WriteDataChannel(p, false)
}
func (w *dataPacketWrapper) Close() error { return w.inner.Close() }
func newDataPacketWrapper(inner datachannel.ReadWriteCloser, kind int) *dataPacketWrapper {
return &dataPacketWrapper{inner: inner, kind: kind}
}

View File

@@ -0,0 +1,217 @@
package wbstream
import (
"context"
"encoding/json"
"fmt"
"net"
"net/http"
"sync"
"sync/atomic"
"time"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing-box/adapter"
"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"
)
const (
reconnectInitialDelay = time.Second
reconnectMaxDelay = 16 * time.Second
)
type WBStreamJoiner struct {
logger logger.ContextLogger
OnConnected func(tunnel.DataTunnel)
dialer N.Dialer
dnsRouter adapter.DNSRouter
PCConfig common.PeerConnectionConfigurer
mu sync.Mutex
session *Session
closed bool
stopCh chan struct{}
stopOnce sync.Once
}
func NewWBStreamJoiner(logger logger.ContextLogger, dialer N.Dialer, dnsRouter adapter.DNSRouter, pcConfig common.PeerConnectionConfigurer) *WBStreamJoiner {
return &WBStreamJoiner{
logger: logger,
dialer: dialer,
dnsRouter: dnsRouter,
PCConfig: pcConfig,
stopCh: make(chan struct{}),
}
}
func (j *WBStreamJoiner) RunWithParams(jsonParams string) {
var params struct {
RoomID string `json:"roomId"`
DisplayName string `json:"displayName"`
TunnelMode string `json:"tunnelMode"`
VP8FPS int `json:"vp8Fps"`
VP8Batch int `json:"vp8Batch"`
DualTrack bool `json:"dualTrack"`
Reliable *bool `json:"reliable"`
}
if err := json.Unmarshal([]byte(jsonParams), &params); err != nil {
j.logger.Error(fmt.Sprintf("wbstream-joiner: failed to parse params: %v", err))
return
}
if params.RoomID == "" {
j.logger.Error("wbstream-joiner: missing roomId")
return
}
if params.DisplayName == "" {
params.DisplayName = "Joiner"
}
reliable := params.Reliable != nil && *params.Reliable
httpClient := j.makeHTTPClient()
j.logger.Info(fmt.Sprintf("wbstream-joiner: room=%s name=%s vp8Fps=%d vp8Batch=%d dualTrack=%v", params.RoomID, params.DisplayName, params.VP8FPS, params.VP8Batch, params.DualTrack))
obf, err := tunnel.NewTunnelObfuscator(tunnel.DeriveSecretFromJoinLink(params.RoomID))
if err != nil {
j.logger.Error(fmt.Sprintf("wbstream-joiner: obfuscator init failed: %v", err))
return
}
j.logger.Debug(fmt.Sprintf("wbstream-joiner: obf key-source=%q localEpoch=0x%08x", params.RoomID, obf.LocalEpoch()))
var settingEngine *webrtc.SettingEngine
if j.PCConfig != nil {
se := webrtc.SettingEngine{}
j.PCConfig.ConfigureSettingEngine(&se)
settingEngine = &se
}
var attempt atomic.Int32
j.logger.Info("wbstream-joiner: connecting")
if err := j.runOnce(httpClient, params.RoomID, params.DisplayName, params.TunnelMode, obf, settingEngine, params.VP8FPS, params.VP8Batch, params.DualTrack, reliable, &attempt); err != nil {
j.logger.Error(fmt.Sprintf("wbstream-joiner: %v", err))
return
}
for {
if j.isClosed() {
j.logger.Info("wbstream-joiner: stopped")
return
}
j.logger.Info("wbstream-joiner: tunnel lost")
if !j.waitBeforeRetry(int(attempt.Load())) {
return
}
attempt.Add(1)
if j.isClosed() {
return
}
j.logger.Info(fmt.Sprintf("wbstream-joiner: reconnect attempt #%d", attempt.Load()))
if err := j.runOnce(httpClient, params.RoomID, params.DisplayName, params.TunnelMode, obf, settingEngine, params.VP8FPS, params.VP8Batch, params.DualTrack, reliable, &attempt); err != nil {
j.logger.Warn(fmt.Sprintf("wbstream-joiner: %v, will retry", err))
}
}
}
func (j *WBStreamJoiner) MarkConfigAcked() {
j.mu.Lock()
sess := j.session
j.mu.Unlock()
if sess != nil {
sess.MarkConfigAcked()
}
}
func (j *WBStreamJoiner) Close() {
j.stopOnce.Do(func() { close(j.stopCh) })
j.mu.Lock()
j.closed = true
sess := j.session
j.session = nil
j.mu.Unlock()
if sess != nil {
sess.Close()
}
}
func (j *WBStreamJoiner) runOnce(httpClient *http.Client, roomID, displayName, tunnelMode string, obf *tunnel.TunnelObfuscator, settingEngine *webrtc.SettingEngine, vp8FPS, vp8Batch int, dualTrack, reliable bool, attempt *atomic.Int32) error {
_, roomToken, _, serverURL, authErr := AuthAndGetToken(httpClient, roomID, displayName)
if authErr != nil {
return fmt.Errorf("auth: %w", authErr)
}
j.logger.Debug(fmt.Sprintf("wbstream-joiner: server=%s", serverURL))
sess := NewSession(SessionConfig{
RoomToken: roomToken,
ServerURL: serverURL,
DisplayName: displayName,
TunnelMode: tunnelMode,
Obfuscator: obf,
Logger: j.logger,
SettingEngine: settingEngine,
Dialer: j.dialer,
DNSRouter: j.dnsRouter,
VP8FPS: vp8FPS,
VP8Batch: vp8Batch,
ScreenShare: dualTrack,
IsJoiner: true,
Reliable: reliable,
})
sess.OnConnected = func(tun tunnel.DataTunnel) {
attempt.Store(0)
j.logger.Info("wbstream-joiner: === TUNNEL CONNECTED ===")
if j.OnConnected != nil {
j.OnConnected(tun)
}
}
j.mu.Lock()
if j.closed {
j.mu.Unlock()
sess.Close()
return nil
}
j.session = sess
j.mu.Unlock()
if err := sess.Start(); err != nil {
j.clearSession(sess)
return fmt.Errorf("session: %w", err)
}
<-sess.Done()
sess.Close()
j.clearSession(sess)
return nil
}
func (j *WBStreamJoiner) waitBeforeRetry(attempt int) bool {
delay := common.BackoffWithJitter(attempt, reconnectInitialDelay, reconnectMaxDelay)
j.logger.Debug(fmt.Sprintf("wbstream-joiner: waiting %s before reconnect", delay))
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-timer.C:
return !j.isClosed()
case <-j.stopCh:
return false
}
}
func (j *WBStreamJoiner) clearSession(sess *Session) {
j.mu.Lock()
if j.session == sess {
j.session = nil
}
j.mu.Unlock()
}
func (j *WBStreamJoiner) isClosed() bool {
j.mu.Lock()
defer j.mu.Unlock()
return j.closed
}
func (j *WBStreamJoiner) makeDialContext() func(ctx context.Context, network, addr string) (net.Conn, error) {
return func(ctx context.Context, network, addr string) (net.Conn, error) {
return j.dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
}
}
func (j *WBStreamJoiner) makeHTTPClient() *http.Client {
transport := &http.Transport{DialContext: j.makeDialContext()}
return &http.Client{Timeout: 60 * time.Second, Transport: transport}
}

View File

@@ -0,0 +1,738 @@
package wbstream
import (
"context"
"fmt"
"net"
"sync"
"time"
"github.com/google/uuid"
"github.com/pion/rtp"
"github.com/pion/rtp/codecs"
"github.com/pion/webrtc/v4"
"github.com/sagernet/sing-box/adapter"
"github.com/sagernet/sing-box/transport/call/common"
"github.com/sagernet/sing-box/transport/call/livekit"
"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 peerEntry struct {
sid string
identity string
firstSeen time.Time
state int32
promoted bool
}
const (
TunnelModeVideo = "video"
TunnelModeDC = "dc"
)
type SessionConfig struct {
RoomToken string
ServerURL string
DisplayName string
TunnelMode string
Obfuscator *tunnel.TunnelObfuscator
Logger logger.ContextLogger
SettingEngine *webrtc.SettingEngine
Dialer N.Dialer
DNSRouter adapter.DNSRouter
VP8FPS int
VP8Batch int
RoomID string
AccessToken string
ReadBuf int
ScreenShare bool
IsJoiner bool
Reliable bool
}
type Session struct {
cfg SessionConfig
lk *livekit.Client
sampleTracks []*webrtc.TrackLocalStaticSample
sampleTransceivers []*webrtc.RTPTransceiver
pubReliableDC *webrtc.DataChannel
pubReliableDCReady bool
subReliableDC *webrtc.DataChannel
vp8tun *tunnel.MultiTrackTunnel
kcptun *tunnel.MultiTrackKCPTunnel
dctun *tunnel.DCTunnel
mu sync.Mutex
tunFired bool
done chan struct{}
peersBySID map[string]peerEntry
kickedSIDs map[string]bool
configAcked chan struct{}
configAckedOnce sync.Once
OnConnected func(tunnel.DataTunnel)
OnPeerRestart func()
OnRemoteCandidate func(target int, candidateOrSDP string)
}
func NewSession(cfg SessionConfig) *Session {
return &Session{
cfg: cfg,
done: make(chan struct{}),
configAcked: make(chan struct{}),
}
}
func (s *Session) MarkConfigAcked() {
s.configAckedOnce.Do(func() {
s.cfg.Logger.Debug("[lk] peer acked vp8 config")
close(s.configAcked)
})
}
func (s *Session) Done() <-chan struct{} { return s.done }
func (s *Session) Start() error {
s.lk = livekit.NewClient(livekit.Config{
ServerURL: s.cfg.ServerURL,
Token: s.cfg.RoomToken,
Origin: Origin,
UserAgent: common.UserAgent,
Logger: s.cfg.Logger,
SettingEngine: s.cfg.SettingEngine,
NetDialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
return s.cfg.Dialer.DialContext(ctx, network, M.ParseSocksaddr(addr))
},
DNSRouter: s.cfg.DNSRouter,
})
s.lk.OnReady = s.onLKReady
s.lk.OnTrack = s.onRemoteTrack
s.lk.OnDataChannel = s.onRemoteDataChannel
s.lk.OnPubConnected = s.startTunnel
if s.cfg.AccessToken != "" && s.cfg.RoomID != "" {
s.lk.OnParticipantUpdate = s.onParticipantUpdate
}
s.lk.OnRemoteCandidate = func(target int, ic webrtc.ICECandidateInit) {
if s.OnRemoteCandidate != nil {
s.OnRemoteCandidate(target, ic.Candidate)
}
}
s.lk.OnRemoteSDP = func(target int, _, sdp string) {
if s.OnRemoteCandidate != nil {
s.OnRemoteCandidate(-1, sdp)
}
}
if err := s.lk.Connect(); err != nil {
return err
}
go s.lk.PingLoop()
go func() {
if err := s.lk.ReadLoop(); err != nil {
s.cfg.Logger.Debug(fmt.Sprintf("[lk] read loop ended: %v", err))
}
s.stopTunnels()
close(s.done)
}()
return nil
}
func (s *Session) AdaptTrackCount(peerCount int) {
if peerCount < 1 {
return
}
pubPC := s.lk.PubPC()
if pubPC == nil {
return
}
s.mu.Lock()
current := len(s.sampleTracks)
s.mu.Unlock()
if peerCount == current {
s.cfg.Logger.Debug(fmt.Sprintf("[lk] adapt-track-count: peer=%d current=%d, no change", peerCount, current))
return
}
if peerCount > current {
s.cfg.Logger.Debug(fmt.Sprintf("[lk] adapt-track-count: scaling publisher tracks %d -> %d", current, peerCount))
for i := current; i < peerCount; i++ {
if !s.addPublisherTrack(pubPC, i) {
return
}
}
} else {
s.cfg.Logger.Debug(fmt.Sprintf("[lk] adapt-track-count: shrinking publisher tracks %d -> %d", current, peerCount))
for i := current; i > peerCount; i-- {
if !s.removePublisherTrack() {
return
}
}
}
offer, err := pubPC.CreateOffer(nil)
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] adapt-track-count: create offer: %v", err))
return
}
if err := pubPC.SetLocalDescription(offer); err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] adapt-track-count: set local offer: %v", err))
return
}
if err := s.lk.SendOffer(offer.SDP); err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] adapt-track-count: send offer: %v", err))
return
}
s.cfg.Logger.Debug(fmt.Sprintf("[lk] adapt-track-count: renegotiation offer sent (%d bytes)", len(offer.SDP)))
}
func (s *Session) Close() {
if s.lk != nil {
s.lk.Close()
}
}
func (s *Session) stopTunnels() {
s.mu.Lock()
vp8 := s.vp8tun
kcptun := s.kcptun
s.mu.Unlock()
if kcptun != nil {
kcptun.Stop()
}
if vp8 != nil {
vp8.Stop()
}
}
func (s *Session) onLKReady() {
pubPC := s.lk.PubPC()
if pubPC == nil {
return
}
camID := "videochannel-" + uuid.New().String()
trackCam, err := webrtc.NewTrackLocalStaticSample(
webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeVP8, ClockRate: 90000},
camID, "tunnel-video-"+uuid.New().String(),
)
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] create local cam track: %v", err))
return
}
tracks := []*webrtc.TrackLocalStaticSample{trackCam}
if s.cfg.ScreenShare {
screenID := "screenchannel-" + uuid.New().String()
trackScreen, err := webrtc.NewTrackLocalStaticSample(
webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeVP8, ClockRate: 90000},
screenID, "tunnel-screen-"+uuid.New().String(),
)
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] create local screen track: %v", err))
return
}
tracks = append(tracks, trackScreen)
}
transceivers := make([]*webrtc.RTPTransceiver, 0, len(tracks))
for _, t := range tracks {
trx, err := pubPC.AddTransceiverFromTrack(t,
webrtc.RTPTransceiverInit{Direction: webrtc.RTPTransceiverDirectionSendonly})
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] add transceiver: %v", err))
return
}
transceivers = append(transceivers, trx)
}
s.mu.Lock()
s.sampleTracks = tracks
s.sampleTransceivers = transceivers
s.mu.Unlock()
ordered := true
dc, err := pubPC.CreateDataChannel("_reliable", &webrtc.DataChannelInit{
Ordered: &ordered,
})
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] create reliable DC: %v", err))
return
}
s.mu.Lock()
s.pubReliableDC = dc
s.mu.Unlock()
dc.OnOpen(func() {
s.cfg.Logger.Debug("[lk] reliable DC open")
s.mu.Lock()
s.pubReliableDCReady = true
s.mu.Unlock()
s.maybeStartDCTunnel()
})
for i, t := range tracks {
source := livekit.TrackSourceCamera
if i > 0 {
source = livekit.TrackSourceScreenShare
}
if err := s.lk.SendAddTrack(t.ID(), "videochannel",
livekit.TrackTypeVideo, source, 1280, 720); err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] send add-track: %v", err))
return
}
}
offer, err := pubPC.CreateOffer(nil)
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] create offer: %v", err))
return
}
if err := pubPC.SetLocalDescription(offer); err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] set local offer: %v", err))
return
}
if err := s.lk.SendOffer(offer.SDP); err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] send offer: %v", err))
return
}
s.cfg.Logger.Debug(fmt.Sprintf("[lk] sent publisher offer (%d bytes)", len(offer.SDP)))
}
func (s *Session) startTunnel() {
s.mu.Lock()
if s.vp8tun != nil || len(s.sampleTracks) == 0 {
s.mu.Unlock()
return
}
subs := make([]*tunnel.VP8DataTunnel, 0, len(s.sampleTracks))
for _, t := range s.sampleTracks {
subs = append(subs, tunnel.NewVP8DataTunnelWithQueue(t, s.cfg.Obfuscator, s.cfg.Logger, tunnel.KCPCarrierQueueDepth))
}
s.vp8tun = tunnel.NewMultiTrackTunnel(subs)
s.vp8tun.SetOnPeerRestart(func() {
s.cfg.Logger.Debug("[wb] peer epoch changed, signalling peer-restart")
s.rearmAutoDetect()
if s.OnPeerRestart != nil {
s.OnPeerRestart()
}
})
s.vp8tun.Start(s.cfg.VP8FPS, s.cfg.VP8Batch)
tun := s.vp8tun
s.mu.Unlock()
s.cfg.Logger.Debug(fmt.Sprintf("[lk] vp8 tunnel writer started tracks=%d", len(subs)))
var active tunnel.DataTunnel = tun
if s.cfg.TunnelMode == TunnelModeVideo && s.cfg.Reliable {
active = s.maybeWrapReliable(tun)
}
if s.cfg.IsJoiner && s.cfg.TunnelMode != TunnelModeDC {
go s.configPingPong(active, len(subs))
}
if s.cfg.TunnelMode == TunnelModeVideo {
s.fireOnConnected(active)
return
}
if s.cfg.TunnelMode == "" {
tun.SetOnData(func(payload []byte) { s.activate(tun, payload) })
}
}
func (s *Session) configPingPong(tun tunnel.DataTunnel, trackCount int) {
frame := tunnel.EncodeVP8Config(s.cfg.VP8FPS, s.cfg.VP8Batch, trackCount)
tun.SendData(frame)
ticker := time.NewTicker(3 * time.Second)
defer ticker.Stop()
for {
select {
case <-s.configAcked:
return
case <-s.done:
return
case <-ticker.C:
s.cfg.Logger.Debug("[lk] resending vp8 config (no ack yet)")
tun.SendData(tunnel.EncodeVP8Config(s.cfg.VP8FPS, s.cfg.VP8Batch, trackCount))
}
}
}
func (s *Session) maybeStartDCTunnel() {
s.mu.Lock()
if s.dctun != nil {
s.mu.Unlock()
return
}
pubDC := s.pubReliableDC
subDC := s.subReliableDC
pubReady := s.pubReliableDCReady
s.mu.Unlock()
if pubDC == nil || subDC == nil || !pubReady {
return
}
if subDC.ReadyState() != webrtc.DataChannelStateOpen {
return
}
subRaw, err := subDC.Detach()
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] detach sub DC: %v", err))
return
}
pubRaw, err := pubDC.Detach()
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] detach pub DC: %v", err))
return
}
readWrapped := newDataPacketWrapper(subRaw, livekit.DataPacketKindReliable)
writeWrapped := newDataPacketWrapper(pubRaw, livekit.DataPacketKindReliable)
readBuf := s.cfg.ReadBuf
if readBuf == 0 {
readBuf = common.DCBufSize
}
dctun := tunnel.NewChunkedDCTunnelFromRaw(readWrapped, writeWrapped, s.cfg.Obfuscator, readBuf, s.cfg.Logger)
if dctun == nil {
return
}
s.mu.Lock()
s.dctun = dctun
s.mu.Unlock()
s.cfg.Logger.Debug("[lk] dc tunnel ready (pub+sub _reliable)")
if s.cfg.TunnelMode == TunnelModeDC {
s.fireOnConnected(dctun)
return
}
if s.cfg.TunnelMode == "" {
dctun.SetOnData(func(payload []byte) { s.activate(dctun, payload) })
}
}
func (s *Session) fireOnConnected(tun tunnel.DataTunnel) {
s.mu.Lock()
if s.tunFired {
s.mu.Unlock()
return
}
s.tunFired = true
s.mu.Unlock()
if s.OnConnected != nil {
s.OnConnected(tun)
}
}
func (s *Session) activate(tun tunnel.DataTunnel, payload []byte) {
s.mu.Lock()
if s.tunFired {
s.mu.Unlock()
return
}
s.tunFired = true
s.mu.Unlock()
var delivered tunnel.DataTunnel = tun
useKCP := false
if _, ok := tun.(*tunnel.MultiTrackTunnel); ok && !tunnel.LooksLikeRelayFrame(payload) {
delivered = s.maybeWrapReliable(tun)
useKCP = true
}
s.cfg.Logger.Debug(fmt.Sprintf("[lk] auto-detected active tunnel: %T", delivered))
if s.OnConnected != nil {
s.OnConnected(delivered)
}
switch v := tun.(type) {
case *tunnel.DCTunnel:
if fwd := v.OnData(); fwd != nil {
fwd(payload)
}
case *tunnel.MultiTrackTunnel:
if useKCP {
if kcptun, ok := delivered.(*tunnel.MultiTrackKCPTunnel); ok {
kcptun.InjectSegment(payload)
}
} else {
v.DeliverData(payload)
}
}
}
func (s *Session) maybeWrapReliable(tun tunnel.DataTunnel) tunnel.DataTunnel {
vp8, ok := tun.(*tunnel.MultiTrackTunnel)
if !ok {
return tun
}
wrapped := tunnel.NewMultiTrackKCPTunnel(vp8, s.cfg.Logger)
s.mu.Lock()
if s.kcptun != nil {
s.kcptun.StopLayer()
}
s.kcptun = wrapped
s.mu.Unlock()
s.cfg.Logger.Debug("[lk] per-track kcp reliability active over video tunnel")
return wrapped
}
func (s *Session) currentVP8Tun() *tunnel.MultiTrackTunnel {
s.mu.Lock()
defer s.mu.Unlock()
return s.vp8tun
}
func (s *Session) removePublisherTrack() bool {
s.mu.Lock()
if len(s.sampleTransceivers) <= 1 || len(s.sampleTracks) <= 1 {
s.mu.Unlock()
s.cfg.Logger.Debug("[lk] adapt-track-count: refusing to remove cam slot")
return false
}
last := len(s.sampleTransceivers) - 1
trx := s.sampleTransceivers[last]
s.sampleTransceivers = s.sampleTransceivers[:last]
s.sampleTracks = s.sampleTracks[:last]
vp8 := s.vp8tun
kcptun := s.kcptun
s.mu.Unlock()
if kcptun != nil {
kcptun.RemoveLastSession()
}
if vp8 != nil {
vp8.RemoveLastSubTunnel()
}
if err := trx.Stop(); err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] adapt-track-count: stop transceiver: %v", err))
return false
}
return true
}
func (s *Session) addPublisherTrack(pubPC *webrtc.PeerConnection, slot int) bool {
labelPrefix := "screenchannel-"
streamPrefix := "tunnel-screen-"
source := livekit.TrackSourceScreenShare
if slot == 0 {
labelPrefix = "videochannel-"
streamPrefix = "tunnel-video-"
source = livekit.TrackSourceCamera
}
track, err := webrtc.NewTrackLocalStaticSample(
webrtc.RTPCodecCapability{MimeType: webrtc.MimeTypeVP8, ClockRate: 90000},
labelPrefix+uuid.New().String(), streamPrefix+uuid.New().String(),
)
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] adapt-track-count: new track slot=%d: %v", slot, err))
return false
}
trx, err := pubPC.AddTransceiverFromTrack(track,
webrtc.RTPTransceiverInit{Direction: webrtc.RTPTransceiverDirectionSendonly})
if err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] adapt-track-count: add transceiver slot=%d: %v", slot, err))
return false
}
if err := s.lk.SendAddTrack(track.ID(), "videochannel",
livekit.TrackTypeVideo, source, 1280, 720); err != nil {
s.cfg.Logger.Error(fmt.Sprintf("[lk] adapt-track-count: send add-track slot=%d: %v", slot, err))
return false
}
s.mu.Lock()
s.sampleTracks = append(s.sampleTracks, track)
s.sampleTransceivers = append(s.sampleTransceivers, trx)
vp8 := s.vp8tun
kcptun := s.kcptun
s.mu.Unlock()
if vp8 != nil {
newSub := tunnel.NewVP8DataTunnelWithQueue(track, s.cfg.Obfuscator, s.cfg.Logger, tunnel.KCPCarrierQueueDepth)
vp8.AddSubTunnel(newSub)
if kcptun != nil {
kcptun.AddSession(newSub)
}
}
return true
}
func (s *Session) rearmAutoDetect() {
if s.cfg.TunnelMode != "" {
return
}
s.mu.Lock()
s.tunFired = false
orphanKCP := s.kcptun
s.kcptun = nil
vp8 := s.vp8tun
dc := s.dctun
s.mu.Unlock()
if orphanKCP != nil {
orphanKCP.StopLayer()
}
if vp8 != nil {
vp8.SetOnData(func(payload []byte) { s.activate(vp8, payload) })
}
if dc != nil {
dc.SetOnData(func(payload []byte) { s.activate(dc, payload) })
}
}
func (s *Session) onRemoteTrack(track *webrtc.TrackRemote, _ *webrtc.RTPReceiver) {
if track.Codec().MimeType != webrtc.MimeTypeVP8 {
go func() {
buf := make([]byte, common.UDPBufSize)
for {
if _, _, err := track.Read(buf); err != nil {
return
}
}
}()
return
}
go s.readVP8Track(track)
}
func (s *Session) readVP8Track(track *webrtc.TrackRemote) {
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 {
s.cfg.Logger.Debug(fmt.Sprintf("[lk-video] recv vp8 frame #%d %d bytes", recvCount, len(frameBuf)))
}
tun := s.currentVP8Tun()
if tun != nil {
tun.HandleFrame(frameBuf)
}
frameBuf = frameBuf[:0]
frameValid = false
}
}
func (s *Session) onRemoteDataChannel(dc *webrtc.DataChannel) {
s.cfg.Logger.Debug(fmt.Sprintf("[lk] remote DC label=%s id=%v", dc.Label(), dc.ID()))
if dc.Label() != "_reliable" {
return
}
s.mu.Lock()
s.subReliableDC = dc
s.mu.Unlock()
dc.OnOpen(func() {
s.cfg.Logger.Debug("[lk] remote _reliable DC open")
s.maybeStartDCTunnel()
})
}
func (s *Session) onParticipantUpdate(updates []livekit.ParticipantInfo) {
selfSID := s.lk.Join().ParticipantSID
s.mu.Lock()
if s.peersBySID == nil {
s.peersBySID = make(map[string]peerEntry)
}
newcomerSIDs := make(map[string]bool)
canPromote := s.cfg.AccessToken != "" && s.cfg.RoomID != ""
for _, p := range updates {
if p.SID == "" || p.SID == selfSID {
continue
}
if p.State == livekit.ParticipantStateDisconnected {
delete(s.peersBySID, p.SID)
delete(s.kickedSIDs, p.SID)
continue
}
if s.kickedSIDs[p.SID] {
continue
}
entry, ok := s.peersBySID[p.SID]
if !ok {
entry = peerEntry{sid: p.SID, identity: p.Identity, firstSeen: time.Now()}
newcomerSIDs[p.SID] = true
}
if p.Identity != "" {
entry.identity = p.Identity
}
entry.state = p.State
s.peersBySID[p.SID] = entry
}
var stale []peerEntry
staleSIDs := make(map[string]bool)
if len(newcomerSIDs) > 0 {
for _, e := range s.peersBySID {
if e.state == livekit.ParticipantStateActive && !newcomerSIDs[e.sid] {
stale = append(stale, e)
staleSIDs[e.sid] = true
}
}
}
var toPromote []peerEntry
if canPromote {
for sid, entry := range s.peersBySID {
if staleSIDs[sid] {
continue
}
if !entry.promoted && entry.state == livekit.ParticipantStateActive && entry.identity != "" {
entry.promoted = true
s.peersBySID[sid] = entry
toPromote = append(toPromote, entry)
}
}
}
s.mu.Unlock()
for _, e := range toPromote {
go s.promotePeer(e.sid, e.identity)
}
if len(stale) == 0 {
return
}
for _, e := range stale {
if e.identity == "" {
continue
}
if err := KickParticipant(common.HttpClient(s.cfg.Dialer), s.cfg.AccessToken, s.cfg.RoomID, e.identity); err != nil {
s.cfg.Logger.Warn(fmt.Sprintf("[wb] kick failed identity=%s: %v", e.identity, err))
continue
}
s.cfg.Logger.Info(fmt.Sprintf("[wb] kicked stale peer identity=%s sid=%s", e.identity, e.sid))
s.mu.Lock()
delete(s.peersBySID, e.sid)
if s.kickedSIDs == nil {
s.kickedSIDs = make(map[string]bool)
}
s.kickedSIDs[e.sid] = true
s.mu.Unlock()
}
}
func (s *Session) promotePeer(sid, identity string) {
if err := SetParticipantPermissions(common.HttpClient(s.cfg.Dialer), s.cfg.AccessToken, s.cfg.RoomID, identity, ModeratorPermissions); err != nil {
s.cfg.Logger.Warn(fmt.Sprintf("[wb] promote failed identity=%s: %v", identity, err))
s.mu.Lock()
if entry, ok := s.peersBySID[sid]; ok {
entry.promoted = false
s.peersBySID[sid] = entry
}
s.mu.Unlock()
return
}
s.cfg.Logger.Info(fmt.Sprintf("[wb] promoted to moderator identity=%s sid=%s", identity, sid))
}

View File

@@ -0,0 +1,293 @@
package wtsignal
import (
"bufio"
"bytes"
"compress/flate"
"context"
"crypto/tls"
"fmt"
"io"
"net"
"net/http"
"net/url"
"sync"
"time"
"github.com/quic-go/quic-go"
"github.com/quic-go/quic-go/http3"
"github.com/quic-go/quic-go/quicvarint"
)
const (
dialTimeout = 15 * time.Second
keepAlivePeriod = 15 * time.Second
maxIdleTimeout = 30 * time.Second
maxMessageSize = 8 << 20
webTransportFrameType uint64 = 0x41
webTransportUniStreamType uint64 = 0x54
settingsEnableWebtransportDraft06 = 0x2b603742
settingsWebTransportEnabled = 0x2c7cf000
settingsWebTransportMaxSessions = 0x14e9cd29
settingsWebTransportMaxSessionsStd = 0xc671706a
closeSessionCapsuleType http3.CapsuleType = 0x2843
protocolHeaderLegacy = "webtransport"
)
type Conn struct {
conn *quic.Conn
stream *quic.Stream
reader *bufio.Reader
compress bool
writeMu sync.Mutex
}
func Dial(endpoint, serverName, resolvedIP string) (*Conn, error) {
target, err := url.Parse(endpoint)
if err != nil {
return nil, err
}
port := target.Port()
if port == "" {
port = "443"
}
compress := target.Query().Get("compression") == "deflate-raw"
tlsConf := &tls.Config{
InsecureSkipVerify: true,
ServerName: serverName,
NextProtos: []string{"h3"},
}
quicConf := &quic.Config{
EnableDatagrams: true,
EnableStreamResetPartialDelivery: true,
KeepAlivePeriod: keepAlivePeriod,
MaxIdleTimeout: maxIdleTimeout,
}
dialCtx, cancel := context.WithTimeout(context.Background(), dialTimeout)
defer cancel()
qconn, err := quic.DialAddrEarly(dialCtx, net.JoinHostPort(resolvedIP, port), tlsConf, quicConf)
if err != nil {
return nil, fmt.Errorf("wt dial: %w", err)
}
// The VK/OK SFU advertises the HTTP/3 datagram setting but does not
// negotiate QUIC transport-level datagrams, which makes quic-go's http3
// layer close the connection. Signaling only uses WebTransport streams, so
// disable HTTP/3 datagrams on our side, and send the draft-06
// ENABLE_WEBTRANSPORT codepoint the SFU expects.
tr := &http3.Transport{
EnableDatagrams: false,
AdditionalSettings: map[uint64]uint64{settingsEnableWebtransportDraft06: 1},
}
control := tr.NewRawClientConn(qconn)
context.AfterFunc(qconn.Context(), func() { tr.Close() })
go acceptStreams(qconn, control)
go acceptUniStreams(qconn, control)
select {
case <-control.ReceivedSettings():
case <-dialCtx.Done():
qconn.CloseWithError(0, "")
return nil, fmt.Errorf("wt settings: %w", dialCtx.Err())
}
settings := control.Settings()
if !settings.EnableExtendedConnect {
qconn.CloseWithError(0, "")
return nil, fmt.Errorf("wt: server did not enable extended connect")
}
if settings.Other[settingsWebTransportEnabled] == 0 &&
settings.Other[settingsEnableWebtransportDraft06] == 0 &&
settings.Other[settingsWebTransportMaxSessions] == 0 &&
settings.Other[settingsWebTransportMaxSessionsStd] == 0 {
qconn.CloseWithError(0, "")
return nil, fmt.Errorf("wt: server did not enable webtransport")
}
requestStr, err := control.OpenRequestStream(dialCtx)
if err != nil {
qconn.CloseWithError(0, "")
return nil, err
}
req := (&http.Request{
Method: http.MethodConnect,
Header: http.Header{},
Proto: protocolHeaderLegacy,
Host: target.Host,
URL: target,
}).WithContext(dialCtx)
if err := requestStr.SendRequestHeader(req); err != nil {
qconn.CloseWithError(0, "")
return nil, err
}
rsp, err := requestStr.ReadResponse()
if err != nil {
qconn.CloseWithError(0, "")
return nil, err
}
if rsp.StatusCode < 200 || rsp.StatusCode >= 300 {
qconn.CloseWithError(0, "")
return nil, fmt.Errorf("wt: connect status %d", rsp.StatusCode)
}
sessionID := uint64(requestStr.StreamID())
go watchSessionClose(requestStr, qconn)
stream, err := qconn.OpenStreamSync(context.Background())
if err != nil {
qconn.CloseWithError(0, "")
return nil, fmt.Errorf("wt open stream: %w", err)
}
streamHdr := quicvarint.Append(nil, webTransportFrameType)
streamHdr = quicvarint.Append(streamHdr, sessionID)
if _, err := stream.Write(streamHdr); err != nil {
qconn.CloseWithError(0, "")
return nil, fmt.Errorf("wt stream header: %w", err)
}
stream.SetReliableBoundary()
return &Conn{
conn: qconn,
stream: stream,
reader: bufio.NewReader(stream),
compress: compress,
}, nil
}
func acceptStreams(qconn *quic.Conn, control *http3.RawClientConn) {
for {
stream, err := qconn.AcceptStream(context.Background())
if err != nil {
return
}
go func() {
typ, err := quicvarint.Peek(stream)
if err != nil {
return
}
if typ != webTransportFrameType {
control.HandleBidirectionalStream(stream)
return
}
if _, err := quicvarint.Read(quicvarint.NewReader(stream)); err != nil {
return
}
if _, err := quicvarint.Read(quicvarint.NewReader(stream)); err != nil {
return
}
io.Copy(io.Discard, stream)
}()
}
}
func acceptUniStreams(qconn *quic.Conn, control *http3.RawClientConn) {
for {
stream, err := qconn.AcceptUniStream(context.Background())
if err != nil {
return
}
go func() {
typ, err := quicvarint.Peek(stream)
if err != nil {
return
}
if typ != webTransportUniStreamType {
control.HandleUnidirectionalStream(stream)
return
}
if _, err := quicvarint.Read(quicvarint.NewReader(stream)); err != nil {
return
}
if _, err := quicvarint.Read(quicvarint.NewReader(stream)); err != nil {
return
}
io.Copy(io.Discard, stream)
}()
}
}
func watchSessionClose(requestStr *http3.RequestStream, qconn *quic.Conn) {
for {
typ, r, err := http3.ParseCapsule(quicvarint.NewReader(requestStr))
if err != nil {
qconn.CloseWithError(0, "")
return
}
if typ == closeSessionCapsuleType {
qconn.CloseWithError(0, "")
return
}
io.Copy(io.Discard, r)
}
}
func (c *Conn) Send(payload []byte) error {
c.writeMu.Lock()
defer c.writeMu.Unlock()
if c.compress {
compressed, err := deflateRaw(payload)
if err != nil {
return err
}
payload = compressed
}
buf := quicvarint.Append(make([]byte, 0, len(payload)+8), uint64(len(payload)))
buf = append(buf, payload...)
_, err := c.stream.Write(buf)
return err
}
func (c *Conn) Recv() ([]byte, error) {
length, err := quicvarint.Read(c.reader)
if err != nil {
return nil, err
}
if length > maxMessageSize {
return nil, fmt.Errorf("wt message too large: %d", length)
}
payload := make([]byte, length)
if _, err := io.ReadFull(c.reader, payload); err != nil {
return nil, err
}
if c.compress {
return inflateRaw(payload)
}
return payload, nil
}
func deflateRaw(payload []byte) ([]byte, error) {
var buf bytes.Buffer
writer, err := flate.NewWriter(&buf, flate.DefaultCompression)
if err != nil {
return nil, err
}
if _, err := writer.Write(payload); err != nil {
return nil, err
}
if err := writer.Close(); err != nil {
return nil, err
}
return buf.Bytes(), nil
}
func inflateRaw(payload []byte) ([]byte, error) {
reader := flate.NewReader(bytes.NewReader(payload))
defer reader.Close()
return io.ReadAll(reader)
}
func (c *Conn) Close() error {
if c.conn != nil {
return c.conn.CloseWithError(0, "")
}
return nil
}