微信小程序\推拉流

This commit is contained in:
2025-12-15 21:57:21 +08:00
parent 2fca6e5c43
commit 73c2073cf1
14 changed files with 2545 additions and 421 deletions

View File

@@ -52,6 +52,8 @@ if errorlevel 1 (
exit /b 1
)
:: Build Linux version with timestamp
echo 1. Building version with timestamp...
go build -o "%DIST_DIR%\%PROJECT%-linux-%TIMESTAMP%"
if %errorlevel% neq 0 (
echo.
@@ -68,6 +70,19 @@ if %errorlevel% neq 0 (
exit /b 1
)
:: Build/Overwrite Linux version without timestamp (simple version)
echo 2. Creating/Overwriting simple version (without timestamp)...
if exist "%DIST_DIR%\%PROJECT%-linux" (
echo Overwriting existing simple version...
del /q "%DIST_DIR%\%PROJECT%-linux" 2>nul
)
copy "%DIST_DIR%\%PROJECT%-linux-%TIMESTAMP%" "%DIST_DIR%\%PROJECT%-linux" >nul
if %errorlevel% neq 0 (
echo ERROR: Failed to create simple version
) else (
echo Simple version created/updated: %PROJECT%-linux
)
echo.
echo [4/5] Building Windows version...
set GOOS=windows
@@ -86,8 +101,9 @@ echo [5/5] Build successful!
echo.
echo Output directory: %cd%\%DIST_DIR%\
echo Generated files:
echo 1. %PROJECT%-linux-%TIMESTAMP%
echo 2. %PROJECT%-windows-%TIMESTAMP%.exe
echo 1. %PROJECT%-linux-%TIMESTAMP% (with timestamp, for archiving)
echo 2. %PROJECT%-linux (simple version, always latest)
echo 3. %PROJECT%-windows-%TIMESTAMP%.exe
echo.
:clean_prompt
@@ -108,9 +124,11 @@ echo.
echo Cleaning old build files...
for %%f in ("%DIST_DIR%\%PROJECT%-*") do (
if not "%%~nxf"=="%PROJECT%-linux-%TIMESTAMP%" (
if not "%%~nxf"=="%PROJECT%-windows-%TIMESTAMP%.exe" (
echo Deleting: %%f
del /q "%%f"
if not "%%~nxf"=="%PROJECT%-linux" (
if not "%%~nxf"=="%PROJECT%-windows-%TIMESTAMP%.exe" (
echo Deleting: %%f
del /q "%%f"
)
)
)
)
@@ -120,7 +138,8 @@ goto :eof
:run_instructions
echo.
echo Run instructions:
echo Linux: .\%DIST_DIR%\%PROJECT%-linux-%TIMESTAMP% --nodeId=node1 --port=12081
echo Linux (simple version): .\%DIST_DIR%\%PROJECT%-linux --nodeId=node1 --port=12081
echo Linux (with timestamp): .\%DIST_DIR%\%PROJECT%-linux-%TIMESTAMP% --nodeId=node1 --port=12081
echo Windows: .\%DIST_DIR%\%PROJECT%-windows-%TIMESTAMP%.exe --nodeId=win1 --port=12080
echo.

View File

@@ -310,6 +310,10 @@ func main() {
}
})
// WebSocket 推流路由(不经过 authGroup避免中间件干扰 WebSocket 升级)
// 路径GET /api/call/ws-push?stream_id=xxx&user_id=xxx&room_id=xxx&token=xxx
r.GET("/api/call/ws-push", api.WSPushHandler)
// 步骤9: 注册HTTP API路由
apiGroup := r.Group("/api")
{

9
go.mod
View File

@@ -3,6 +3,7 @@ module xk-websocket-v2
go 1.24.1
require (
github.com/at-wat/ebml-go v0.17.2
github.com/gin-gonic/gin v1.11.0
github.com/go-redis/redis/v8 v8.11.5
github.com/golang-jwt/jwt/v5 v5.3.0
@@ -13,8 +14,6 @@ require (
github.com/pion/turn/v2 v2.1.6
github.com/pion/webrtc/v3 v3.3.6
github.com/spf13/viper v1.21.0
github.com/yutopp/go-flv v0.3.1
github.com/yutopp/go-rtmp v0.0.7
golang.org/x/crypto v0.45.0
gorm.io/driver/mysql v1.6.0
gorm.io/gorm v1.31.1
@@ -40,15 +39,12 @@ require (
github.com/goccy/go-json v0.10.5 // indirect
github.com/goccy/go-yaml v1.19.0 // indirect
github.com/google/uuid v1.3.1 // indirect
github.com/hashicorp/errwrap v1.1.0 // indirect
github.com/hashicorp/go-multierror v1.1.0 // indirect
github.com/jinzhu/inflection v1.0.0 // indirect
github.com/jinzhu/now v1.1.5 // indirect
github.com/json-iterator/go v1.1.12 // indirect
github.com/klauspost/cpuid/v2 v2.3.0 // indirect
github.com/leodido/go-urn v1.4.0 // indirect
github.com/mattn/go-isatty v0.0.20 // indirect
github.com/mitchellh/mapstructure v1.4.1 // indirect
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
github.com/modern-go/reflect2 v1.0.2 // indirect
github.com/pelletier/go-toml/v2 v2.2.4 // indirect
@@ -65,12 +61,10 @@ require (
github.com/pion/stun v0.6.1 // indirect
github.com/pion/transport/v2 v2.2.10 // indirect
github.com/pion/transport/v3 v3.1.1 // indirect
github.com/pkg/errors v0.9.1 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect
github.com/quic-go/qpack v0.6.0 // indirect
github.com/quic-go/quic-go v0.57.1 // indirect
github.com/sagikazarmark/locafero v0.12.0 // indirect
github.com/sirupsen/logrus v1.7.0 // indirect
github.com/spf13/afero v1.15.0 // indirect
github.com/spf13/cast v1.10.0 // indirect
github.com/spf13/pflag v1.0.10 // indirect
@@ -79,7 +73,6 @@ require (
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
github.com/ugorji/go/codec v1.3.1 // indirect
github.com/wlynxg/anet v0.0.5 // indirect
github.com/yutopp/go-amf0 v0.1.0 // indirect
go.uber.org/mock v0.6.0 // indirect
go.yaml.in/yaml/v3 v3.0.4 // indirect
golang.org/x/arch v0.23.0 // indirect

25
go.sum
View File

@@ -1,5 +1,7 @@
filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA=
filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4=
github.com/at-wat/ebml-go v0.17.2 h1:FJ89W5V6jDklwBNEPQoFKnQnBxbSpXB72qrobW3iIWY=
github.com/at-wat/ebml-go v0.17.2/go.mod h1:w1cJs7zmGsb5nnSvhWGKLCxvfu4FVx5ERvYDIalj1ww=
github.com/bytedance/gopkg v0.1.3 h1:TPBSwH8RsouGCBcMBktLt1AymVo2TVsBVCY4b6TnZ/M=
github.com/bytedance/gopkg v0.1.3/go.mod h1:576VvJ+eJgyCzdjS+c4+77QF3p7ubbtiKARP3TxducM=
github.com/bytedance/sonic v1.14.2 h1:k1twIoe97C1DtYUo+fZQy865IuHia4PR5RPiuGPPIIE=
@@ -15,8 +17,6 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78=
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc=
github.com/fortytw2/leaktest v1.2.0 h1:cj6GCiwJDH7l3tMHLjZDo0QqPtrXJiWSI9JgpeQKw+Q=
github.com/fortytw2/leaktest v1.2.0/go.mod h1:jDsjWgpAGjm2CA7WthBh/CdZYEPF31XHquHwclZch5g=
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k=
@@ -54,11 +54,6 @@ github.com/google/uuid v1.3.1 h1:KjJaJ9iWZ3jOFZIf1Lqf4laDRCasjl0BCmnEGxkdLb4=
github.com/google/uuid v1.3.1/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
github.com/hashicorp/errwrap v1.0.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4=
github.com/hashicorp/errwrap v1.1.0 h1:OxrOeh75EUXMY8TBjag2fzXGZ40LB6IKw45YeGUDY2I=
github.com/hashicorp/errwrap v1.1.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4=
github.com/hashicorp/go-multierror v1.1.0 h1:B9UzwGQJehnUY1yNrnwREHc3fGbC2xefo8g4TbElacI=
github.com/hashicorp/go-multierror v1.1.0/go.mod h1:spPvp8C1qA32ftKqdAHm4hHTbPw+vmowP0z+KUhOZdA=
github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD/E=
github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc=
github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ=
@@ -78,8 +73,6 @@ github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ=
github.com/leodido/go-urn v1.4.0/go.mod h1:bvxc+MVxLKB4z00jd1z+Dvzr47oO32F/QSNjSBOlFxI=
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
github.com/mitchellh/mapstructure v1.4.1 h1:CpVNEelQCZBooIPDn+AR3NpivK/TIKU8bDxdASFVQag=
github.com/mitchellh/mapstructure v1.4.1/go.mod h1:bFUtVrKA4DC2yAKiSyO/QUcy7e+RRV2QTWOzhPopBRo=
github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg=
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q=
@@ -138,8 +131,6 @@ github.com/pion/turn/v2 v2.1.6 h1:Xr2niVsiPTB0FPtt+yAWKFUkU1eotQbGgpTIld4x1Gc=
github.com/pion/turn/v2 v2.1.6/go.mod h1:huEpByKKHix2/b9kmTAM3YoX6MKP+/D//0ClgUYR2fY=
github.com/pion/webrtc/v3 v3.3.6 h1:7XAh4RPtlY1Vul6/GmZrv7z+NnxKA6If0KStXBI2ZLE=
github.com/pion/webrtc/v3 v3.3.6/go.mod h1:zyN7th4mZpV27eXybfR/cnUf3J2DRy8zw/mdjD9JTNM=
github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/quic-go/qpack v0.6.0 h1:g7W+BMYynC1LbYLSqRt8PBg5Tgwxn214ZZR34VIOjz8=
@@ -150,8 +141,6 @@ github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjR
github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog=
github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88eegjfxfHb4=
github.com/sagikazarmark/locafero v0.12.0/go.mod h1:sZh36u/YSZ918v0Io+U9ogLYQJ9tLLBmM4eneO6WwsI=
github.com/sirupsen/logrus v1.7.0 h1:ShrD1U9pZB12TX0cVy0DtePoCH97K8EtX+mg7ZARUtM=
github.com/sirupsen/logrus v1.7.0/go.mod h1:yWOB1SBYBC5VeMP7gHvWumXLIWorT60ONWic61uBYv0=
github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I=
github.com/spf13/afero v1.15.0/go.mod h1:NC2ByUVxtQs4b3sIUphxK0NioZnmxgyCrfzeuq8lxMg=
github.com/spf13/cast v1.10.0 h1:h2x0u2shc1QuLHfxi+cTJvs30+ZAHOGRic8uyGTDWxY=
@@ -164,11 +153,9 @@ github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+
github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw=
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXfy6kDkUVs=
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
github.com/stretchr/testify v1.8.3/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
@@ -185,12 +172,6 @@ github.com/wlynxg/anet v0.0.3/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU=
github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA=
github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY=
github.com/yutopp/go-amf0 v0.1.0 h1:a3UeBZG7nRF0zfvmPn2iAfNo1RGzUpHz1VyJD2oGrik=
github.com/yutopp/go-amf0 v0.1.0/go.mod h1:QzDOBr9RV6sQh6E5GFEJROZbU0iQKijORBmprkb3FIk=
github.com/yutopp/go-flv v0.3.1 h1:4ILK6OgCJgUNm2WOjaucWM5lUHE0+sLNPdjq3L0Xtjk=
github.com/yutopp/go-flv v0.3.1/go.mod h1:pAlHPSVRMv5aCUKmGOS/dZn/ooTgnc09qOPmiUNMubs=
github.com/yutopp/go-rtmp v0.0.7 h1:sKKm1MVV3ANbJHZlf3Kq8ecq99y5U7XnDUDxSjuK7KU=
github.com/yutopp/go-rtmp v0.0.7/go.mod h1:KSwrC9Xj5Kf18EUlk1g7CScecjXfIqc0J5q+S0u6Irc=
go.uber.org/mock v0.6.0 h1:hyF9dfmbgIX5EfOdasqLsWD6xqpNZlXblLB/Dbnwv3Y=
go.uber.org/mock v0.6.0/go.mod h1:KiVJ4BqZJaMj4svdfmHM0AUx4NJYO8ZNpPnZn1Z+BBU=
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
@@ -222,12 +203,10 @@ golang.org/x/sync v0.1.0/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.18.0 h1:kr88TuHDroi+UVf+0hZnirlk8o8T+4MrK6mr60WkH/I=
golang.org/x/sync v0.18.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20191026070338-33540a1f6037/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.3.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.5.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.7.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=

View File

@@ -11,6 +11,7 @@ package api
import (
"fmt"
"log"
"xk-websocket-v2/internal/mediaserver"
"xk-websocket-v2/internal/utils"
@@ -59,10 +60,13 @@ type WebRTCICERequest struct {
type JoinCallRoomResponse struct {
RoomID string `json:"room_id"`
Platform string `json:"platform"`
ICEServers []ICEServerConfig `json:"ice_servers,omitempty"` // H5/App 用
ICEServers []ICEServerConfig `json:"ice_servers,omitempty"` // H5/App 用 (WebRTC 模式)
WSPushURL string `json:"ws_push_url,omitempty"` // H5/App 用 (RTMP 模式 WebSocket 推流地址)
SelfFLVURL string `json:"self_flv_url,omitempty"` // H5/App 用 - 自己的 HTTP-FLV 地址(供小程序拉流)
FlvPullURLs []PullURLInfo `json:"flv_pull_urls,omitempty"` // H5/App 拉取小程序流的 FLV 地址
PushURL string `json:"push_url,omitempty"` // 小程序用
PullURLs []PullURLInfo `json:"pull_urls,omitempty"` // 小程序用
PushURL string `json:"push_url,omitempty"` // 小程序用 RTMP 推流地址
FLVURL string `json:"flv_url,omitempty"` // 小程序用 - 自己的 HTTP-FLV 地址(供其他端拉流)
PullURLs []PullURLInfo `json:"pull_urls,omitempty"` // 小程序用 RTMP 拉流地址
Participants []ParticipantInfo `json:"participants"`
}
@@ -154,17 +158,40 @@ func JoinCallRoomHandler(c *gin.Context) {
switch req.Platform {
case "h5", "app":
// H5/App 使用 WebRTC
// H5/App 使用 WebRTC 或 RTMP 模式
pType := mediaserver.ParticipantTypeWebRTC
_, err := room.AddParticipant(req.UserID, pType)
participant, err := room.AddParticipant(req.UserID, pType)
if err != nil {
utils.BadRequest(c, err.Error())
return
}
// 返回 ICE 服务器配置
// 返回 ICE 服务器配置 (WebRTC 模式)
response.ICEServers = getICEServers(req.UserID)
// 生成 WebSocket 推流地址 (RTMP 模式)
// 格式: ws://host:port/api/call/ws-push?stream_id=xxx&user_id=xxx&room_id=xxx&token=xxx
rtmpServer := ms.GetRTMP()
if rtmpServer != nil {
// 为 H5/App 生成流信息
stream, err := rtmpServer.GenerateStreamURLs(req.RoomID, req.UserID)
if err == nil {
// 更新参与者的流地址
room.SetParticipantRTMPURLs(req.UserID, stream.PushURL, stream.PullURL, stream.FLVURL, stream.ID)
participant.PushURL = stream.PushURL
participant.PullURL = stream.PullURL
participant.FLVURL = stream.FLVURL
participant.StreamID = stream.ID
// 构建 WebSocket 推流 URL
token := mediaserver.GenerateStreamToken(stream.ID)
response.WSPushURL = fmt.Sprintf("wss://g-ws.nailaoyun.cn/api/call/ws-push?stream_id=%s&user_id=%s&room_id=%s&token=%s",
stream.ID, req.UserID, req.RoomID, token)
// 返回 H5/App 自己的 FLV 地址(供小程序拉流)
response.SelfFLVURL = stream.FLVURL
}
}
// 获取房间内小程序用户的 FLV 拉流地址(用于 Web 端播放小程序流)
flvPullURLs := make([]PullURLInfo, 0)
for _, p := range room.GetOtherParticipants(req.UserID) {
@@ -209,6 +236,7 @@ func JoinCallRoomHandler(c *gin.Context) {
participant.StreamID = stream.ID
response.PushURL = stream.PushURL
response.FLVURL = stream.FLVURL // 返回自己的 FLV 地址,供客户端发送给其他端
// 获取房间内其他用户的拉流地址
pullURLs := make([]PullURLInfo, 0)
@@ -432,5 +460,33 @@ func RegisterCallRoutes(router *gin.RouterGroup) {
callGroup.POST("/offer", WebRTCOfferHandler)
callGroup.POST("/ice", WebRTCICEHandler)
callGroup.GET("/ice-servers", GetICEServersHandler)
// 注意ws-push 路由已移到全局main.go避免中间件干扰 WebSocket 升级
}
}
// WSPushHandler WebSocket 推流处理
// GET /api/call/ws-push?stream_id=xxx&user_id=xxx&room_id=xxx&token=xxx
func WSPushHandler(c *gin.Context) {
// 详细日志:确认请求到达
log.Printf("🔌 [WSPush] 收到请求: %s from %s", c.Request.URL.String(), c.ClientIP())
log.Printf("🔌 [WSPush] Headers: Upgrade=%s Connection=%s Origin=%s",
c.GetHeader("Upgrade"), c.GetHeader("Connection"), c.GetHeader("Origin"))
ms := mediaserver.GetServer()
if ms.GetConfig() == nil || !ms.GetConfig().Enabled {
log.Printf("❌ [WSPush] 媒体服务未启用")
utils.Error(c, 503, "媒体服务未启用")
return
}
wsProxy := ms.GetWSProxy()
if wsProxy == nil {
log.Printf("❌ [WSPush] WebSocket代理未启用")
utils.Error(c, 503, "WebSocket代理未启用")
return
}
log.Printf("✅ [WSPush] 转交给 WebSocket 代理处理")
// 交给 WebSocket 代理处理
wsProxy.HandleWebSocket(c.Writer, c.Request)
}

View File

@@ -107,6 +107,15 @@ func (r *Room) AddParticipant(userID string, pType ParticipantType) (*Participan
// 检查是否已存在
if p, exists := r.Participants[userID]; exists {
// 用户重新加入,清除旧的流信息以便重新生成
oldStreamID := p.StreamID
p.PushURL = ""
p.PullURL = ""
p.FLVURL = ""
p.StreamID = ""
p.Type = pType
p.JoinedAt = time.Now()
log.Printf("🔄 [Room:%s] 用户 %s 重新加入,已清除旧流信息 (旧StreamID: %s)", r.ID, userID, oldStreamID)
return p, nil
}

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,427 @@
/**
* package mediaserver
*
* AMF0 (Action Message Format 0) 编解码实现
* 用于 RTMP 命令消息的序列化和反序列化
*/
package mediaserver
import (
"bytes"
"encoding/binary"
"fmt"
"math"
)
// AMF0 数据类型标记
const (
AMF0_NUMBER = 0x00 // 8 bytes double
AMF0_BOOLEAN = 0x01 // 1 byte
AMF0_STRING = 0x02 // 2 bytes length + data
AMF0_OBJECT = 0x03 // key-value pairs
AMF0_MOVIECLIP = 0x04 // reserved
AMF0_NULL = 0x05 // no data
AMF0_UNDEFINED = 0x06 // no data
AMF0_REFERENCE = 0x07 // 2 bytes
AMF0_ECMA_ARRAY = 0x08 // associative array
AMF0_OBJECT_END = 0x09 // object end marker
AMF0_STRICT_ARRAY = 0x0A // strict array
AMF0_DATE = 0x0B // 8 bytes double + 2 bytes timezone
AMF0_LONG_STRING = 0x0C // 4 bytes length + data
AMF0_UNSUPPORTED = 0x0D
AMF0_RECORDSET = 0x0E // reserved
AMF0_XML_DOCUMENT = 0x0F
AMF0_TYPED_OBJECT = 0x10
AMF0_AVMPLUS = 0x11 // switch to AMF3
)
// AMF0Object 表示一个 AMF0 对象
type AMF0Object map[string]interface{}
// AMF0Encoder AMF0 编码器
type AMF0Encoder struct {
buf *bytes.Buffer
}
// AMF0Decoder AMF0 解码器
type AMF0Decoder struct {
data []byte
offset int
}
// NewAMF0Encoder 创建 AMF0 编码器
func NewAMF0Encoder() *AMF0Encoder {
return &AMF0Encoder{
buf: new(bytes.Buffer),
}
}
// NewAMF0Decoder 创建 AMF0 解码器
func NewAMF0Decoder(data []byte) *AMF0Decoder {
return &AMF0Decoder{
data: data,
offset: 0,
}
}
// Bytes 获取编码后的字节
func (e *AMF0Encoder) Bytes() []byte {
return e.buf.Bytes()
}
// Reset 重置编码器
func (e *AMF0Encoder) Reset() {
e.buf.Reset()
}
// EncodeNumber 编码数字
func (e *AMF0Encoder) EncodeNumber(val float64) {
e.buf.WriteByte(AMF0_NUMBER)
bits := math.Float64bits(val)
binary.Write(e.buf, binary.BigEndian, bits)
}
// EncodeBoolean 编码布尔值
func (e *AMF0Encoder) EncodeBoolean(val bool) {
e.buf.WriteByte(AMF0_BOOLEAN)
if val {
e.buf.WriteByte(1)
} else {
e.buf.WriteByte(0)
}
}
// EncodeString 编码字符串
func (e *AMF0Encoder) EncodeString(val string) {
data := []byte(val)
if len(data) > 0xFFFF {
// Long string
e.buf.WriteByte(AMF0_LONG_STRING)
binary.Write(e.buf, binary.BigEndian, uint32(len(data)))
} else {
e.buf.WriteByte(AMF0_STRING)
binary.Write(e.buf, binary.BigEndian, uint16(len(data)))
}
e.buf.Write(data)
}
// EncodeNull 编码 null
func (e *AMF0Encoder) EncodeNull() {
e.buf.WriteByte(AMF0_NULL)
}
// EncodeObject 编码对象
func (e *AMF0Encoder) EncodeObject(obj AMF0Object) {
e.buf.WriteByte(AMF0_OBJECT)
for key, val := range obj {
// 写入属性名(不带类型标记)
binary.Write(e.buf, binary.BigEndian, uint16(len(key)))
e.buf.WriteString(key)
// 写入值
e.EncodeValue(val)
}
// 写入对象结束标记
e.buf.Write([]byte{0, 0, AMF0_OBJECT_END})
}
// EncodeEcmaArray 编码 ECMA 数组
func (e *AMF0Encoder) EncodeEcmaArray(obj AMF0Object) {
e.buf.WriteByte(AMF0_ECMA_ARRAY)
binary.Write(e.buf, binary.BigEndian, uint32(len(obj)))
for key, val := range obj {
binary.Write(e.buf, binary.BigEndian, uint16(len(key)))
e.buf.WriteString(key)
e.EncodeValue(val)
}
e.buf.Write([]byte{0, 0, AMF0_OBJECT_END})
}
// EncodeValue 编码任意值
func (e *AMF0Encoder) EncodeValue(val interface{}) {
switch v := val.(type) {
case float64:
e.EncodeNumber(v)
case float32:
e.EncodeNumber(float64(v))
case int:
e.EncodeNumber(float64(v))
case int64:
e.EncodeNumber(float64(v))
case int32:
e.EncodeNumber(float64(v))
case uint32:
e.EncodeNumber(float64(v))
case bool:
e.EncodeBoolean(v)
case string:
e.EncodeString(v)
case nil:
e.EncodeNull()
case AMF0Object:
e.EncodeObject(v)
case map[string]interface{}:
e.EncodeObject(AMF0Object(v))
default:
e.EncodeNull()
}
}
// Remaining 返回剩余未解析的字节数
func (d *AMF0Decoder) Remaining() int {
return len(d.data) - d.offset
}
// DecodeAll 解码所有值
func (d *AMF0Decoder) DecodeAll() ([]interface{}, error) {
var values []interface{}
for d.Remaining() > 0 {
val, err := d.DecodeValue()
if err != nil {
break
}
values = append(values, val)
}
return values, nil
}
// DecodeValue 解码一个值
func (d *AMF0Decoder) DecodeValue() (interface{}, error) {
if d.Remaining() < 1 {
return nil, fmt.Errorf("数据不足")
}
marker := d.data[d.offset]
d.offset++
switch marker {
case AMF0_NUMBER:
return d.decodeNumber()
case AMF0_BOOLEAN:
return d.decodeBoolean()
case AMF0_STRING:
return d.decodeString()
case AMF0_OBJECT:
return d.decodeObject()
case AMF0_NULL, AMF0_UNDEFINED:
return nil, nil
case AMF0_ECMA_ARRAY:
return d.decodeEcmaArray()
case AMF0_STRICT_ARRAY:
return d.decodeStrictArray()
case AMF0_LONG_STRING:
return d.decodeLongString()
case AMF0_DATE:
return d.decodeDate()
default:
return nil, fmt.Errorf("未知的 AMF0 类型: 0x%02X", marker)
}
}
// decodeNumber 解码数字
func (d *AMF0Decoder) decodeNumber() (float64, error) {
if d.Remaining() < 8 {
return 0, fmt.Errorf("数据不足以解码 number")
}
bits := binary.BigEndian.Uint64(d.data[d.offset : d.offset+8])
d.offset += 8
return math.Float64frombits(bits), nil
}
// decodeBoolean 解码布尔值
func (d *AMF0Decoder) decodeBoolean() (bool, error) {
if d.Remaining() < 1 {
return false, fmt.Errorf("数据不足以解码 boolean")
}
val := d.data[d.offset] != 0
d.offset++
return val, nil
}
// decodeString 解码字符串
func (d *AMF0Decoder) decodeString() (string, error) {
if d.Remaining() < 2 {
return "", fmt.Errorf("数据不足以解码 string length")
}
length := int(binary.BigEndian.Uint16(d.data[d.offset : d.offset+2]))
d.offset += 2
if d.Remaining() < length {
return "", fmt.Errorf("数据不足以解码 string data")
}
str := string(d.data[d.offset : d.offset+length])
d.offset += length
return str, nil
}
// decodeLongString 解码长字符串
func (d *AMF0Decoder) decodeLongString() (string, error) {
if d.Remaining() < 4 {
return "", fmt.Errorf("数据不足以解码 long string length")
}
length := int(binary.BigEndian.Uint32(d.data[d.offset : d.offset+4]))
d.offset += 4
if d.Remaining() < length {
return "", fmt.Errorf("数据不足以解码 long string data")
}
str := string(d.data[d.offset : d.offset+length])
d.offset += length
return str, nil
}
// decodeObject 解码对象
func (d *AMF0Decoder) decodeObject() (AMF0Object, error) {
obj := make(AMF0Object)
for {
// 读取属性名长度
if d.Remaining() < 2 {
return nil, fmt.Errorf("数据不足以解码对象属性名长度")
}
nameLen := int(binary.BigEndian.Uint16(d.data[d.offset : d.offset+2]))
d.offset += 2
// 检查是否到达对象结束
if nameLen == 0 {
if d.Remaining() < 1 || d.data[d.offset] != AMF0_OBJECT_END {
return nil, fmt.Errorf("对象结束标记缺失")
}
d.offset++
break
}
// 读取属性名
if d.Remaining() < nameLen {
return nil, fmt.Errorf("数据不足以解码对象属性名")
}
name := string(d.data[d.offset : d.offset+nameLen])
d.offset += nameLen
// 读取属性值
val, err := d.DecodeValue()
if err != nil {
return nil, fmt.Errorf("解码对象属性值失败: %w", err)
}
obj[name] = val
}
return obj, nil
}
// decodeEcmaArray 解码 ECMA 数组
func (d *AMF0Decoder) decodeEcmaArray() (AMF0Object, error) {
if d.Remaining() < 4 {
return nil, fmt.Errorf("数据不足以解码 ECMA 数组长度")
}
// 读取数组长度(但实际不使用,因为以 object end 结束)
d.offset += 4
return d.decodeObject()
}
// decodeStrictArray 解码严格数组
func (d *AMF0Decoder) decodeStrictArray() ([]interface{}, error) {
if d.Remaining() < 4 {
return nil, fmt.Errorf("数据不足以解码严格数组长度")
}
length := int(binary.BigEndian.Uint32(d.data[d.offset : d.offset+4]))
d.offset += 4
arr := make([]interface{}, length)
for i := 0; i < length; i++ {
val, err := d.DecodeValue()
if err != nil {
return nil, fmt.Errorf("解码数组元素失败: %w", err)
}
arr[i] = val
}
return arr, nil
}
// decodeDate 解码日期
func (d *AMF0Decoder) decodeDate() (float64, error) {
if d.Remaining() < 10 {
return 0, fmt.Errorf("数据不足以解码 date")
}
bits := binary.BigEndian.Uint64(d.data[d.offset : d.offset+8])
d.offset += 8
// 跳过时区信息 (2 bytes)
d.offset += 2
return math.Float64frombits(bits), nil
}
// EncodeAMF0 编码多个值为 AMF0 格式
func EncodeAMF0(values ...interface{}) []byte {
encoder := NewAMF0Encoder()
for _, val := range values {
encoder.EncodeValue(val)
}
return encoder.Bytes()
}
// DecodeAMF0 从 AMF0 格式解码多个值
func DecodeAMF0(data []byte) ([]interface{}, error) {
decoder := NewAMF0Decoder(data)
return decoder.DecodeAll()
}
// EncodeConnectResult 编码 connect 响应
func EncodeConnectResult(transactionID float64) []byte {
encoder := NewAMF0Encoder()
// _result
encoder.EncodeString("_result")
encoder.EncodeNumber(transactionID)
// Properties
encoder.EncodeObject(AMF0Object{
"fmsVer": "FMS/3,5,7,7009",
"capabilities": float64(31),
"mode": float64(1),
})
// Information
encoder.EncodeObject(AMF0Object{
"level": "status",
"code": "NetConnection.Connect.Success",
"description": "Connection succeeded.",
"objectEncoding": float64(0),
})
return encoder.Bytes()
}
// EncodeCreateStreamResult 编码 createStream 响应
func EncodeCreateStreamResult(transactionID float64, streamID float64) []byte {
encoder := NewAMF0Encoder()
encoder.EncodeString("_result")
encoder.EncodeNumber(transactionID)
encoder.EncodeNull()
encoder.EncodeNumber(streamID)
return encoder.Bytes()
}
// EncodeOnStatus 编码 onStatus 消息
func EncodeOnStatus(code, level, description string) []byte {
encoder := NewAMF0Encoder()
encoder.EncodeString("onStatus")
encoder.EncodeNumber(0)
encoder.EncodeNull()
encoder.EncodeObject(AMF0Object{
"level": level,
"code": code,
"description": description,
})
return encoder.Bytes()
}
// EncodeOnBWDone 编码 onBWDone 消息
func EncodeOnBWDone() []byte {
encoder := NewAMF0Encoder()
encoder.EncodeString("onBWDone")
encoder.EncodeNumber(0)
encoder.EncodeNull()
return encoder.Bytes()
}

View File

@@ -0,0 +1,161 @@
/**
* package mediaserver
*
* RTMP 握手实现
* 支持 Simple Handshake (Version 3)
*/
package mediaserver
import (
"bytes"
"crypto/rand"
"encoding/binary"
"fmt"
"io"
"log"
"net"
"time"
)
// 握手相关常量
const (
HANDSHAKE_SIZE = 1536
RTMP_VERSION = 3
)
// DoHandshake 执行 RTMP 握手(服务器端)
// 握手流程:
// 1. 接收 C0 + C1 (1 + 1536 bytes)
// 2. 发送 S0 + S1 + S2 (1 + 1536 + 1536 bytes)
// 3. 接收 C2 (1536 bytes)
func DoHandshake(conn net.Conn, timeout time.Duration) error {
// 设置超时
conn.SetDeadline(time.Now().Add(timeout))
defer conn.SetDeadline(time.Time{}) // 清除超时
// 1. 接收 C0 (1 byte: version)
c0 := make([]byte, 1)
if _, err := io.ReadFull(conn, c0); err != nil {
return fmt.Errorf("读取 C0 失败: %w", err)
}
version := c0[0]
if version != RTMP_VERSION {
// 尝试兼容其他版本
log.Printf("⚠️ [RTMP Handshake] 客户端版本: %d (期望: %d)", version, RTMP_VERSION)
}
// 2. 接收 C1 (1536 bytes)
c1 := make([]byte, HANDSHAKE_SIZE)
if _, err := io.ReadFull(conn, c1); err != nil {
return fmt.Errorf("读取 C1 失败: %w", err)
}
// C1 结构:
// - time (4 bytes): 时间戳
// - zero (4 bytes): 必须为0简单握手或版本信息复杂握手
// - random (1528 bytes): 随机数据
c1Time := binary.BigEndian.Uint32(c1[0:4])
c1Zero := binary.BigEndian.Uint32(c1[4:8])
log.Printf("📡 [RTMP Handshake] C1: time=%d zero=%d", c1Time, c1Zero)
// 3. 生成 S0 + S1 + S2
s0 := []byte{RTMP_VERSION}
// S1 (1536 bytes): time(4) + zero(4) + random(1528)
s1 := make([]byte, HANDSHAKE_SIZE)
binary.BigEndian.PutUint32(s1[0:4], uint32(time.Now().Unix())) // time
binary.BigEndian.PutUint32(s1[4:8], 0) // zero
rand.Read(s1[8:]) // random
// S2 (1536 bytes): 回显 C1简单握手
// - time (4 bytes): C1 的时间戳
// - time2 (4 bytes): S1 的时间戳
// - random echo (1528 bytes): C1 的随机数据
s2 := make([]byte, HANDSHAKE_SIZE)
copy(s2[0:4], c1[0:4]) // 回显 C1 时间戳
binary.BigEndian.PutUint32(s2[4:8], binary.BigEndian.Uint32(s1[0:4])) // S1 时间戳
copy(s2[8:], c1[8:]) // 回显 C1 随机数据
// 4. 发送 S0 + S1 + S2
response := make([]byte, 0, 1+HANDSHAKE_SIZE*2)
response = append(response, s0...)
response = append(response, s1...)
response = append(response, s2...)
if _, err := conn.Write(response); err != nil {
return fmt.Errorf("发送 S0+S1+S2 失败: %w", err)
}
// 5. 接收 C2 (1536 bytes)
c2 := make([]byte, HANDSHAKE_SIZE)
if _, err := io.ReadFull(conn, c2); err != nil {
return fmt.Errorf("读取 C2 失败: %w", err)
}
// 验证 C2 (可选,简单握手可以跳过)
// C2 应该回显 S1 的数据
c2Time := binary.BigEndian.Uint32(c2[0:4])
if c2Time != binary.BigEndian.Uint32(s1[0:4]) {
log.Printf("⚠️ [RTMP Handshake] C2 时间戳不匹配 (收到: %d, 期望: %d)", c2Time, binary.BigEndian.Uint32(s1[0:4]))
// 不返回错误,继续处理
}
// 简单验证:比较随机数据的前几个字节
if !bytes.Equal(c2[8:16], s1[8:16]) {
log.Printf("⚠️ [RTMP Handshake] C2 随机数据不匹配")
// 不返回错误,继续处理
}
log.Printf("✅ [RTMP Handshake] 握手完成 | remote=%s", conn.RemoteAddr())
return nil
}
// DoClientHandshake 执行 RTMP 握手(客户端)
// 用于测试或代理场景
func DoClientHandshake(conn net.Conn, timeout time.Duration) error {
// 设置超时
conn.SetDeadline(time.Now().Add(timeout))
defer conn.SetDeadline(time.Time{})
// 1. 发送 C0 + C1
c0 := []byte{RTMP_VERSION}
c1 := make([]byte, HANDSHAKE_SIZE)
binary.BigEndian.PutUint32(c1[0:4], uint32(time.Now().Unix())) // time
binary.BigEndian.PutUint32(c1[4:8], 0) // zero
rand.Read(c1[8:]) // random
if _, err := conn.Write(append(c0, c1...)); err != nil {
return fmt.Errorf("发送 C0+C1 失败: %w", err)
}
// 2. 接收 S0 + S1 + S2
s0s1s2 := make([]byte, 1+HANDSHAKE_SIZE*2)
if _, err := io.ReadFull(conn, s0s1s2); err != nil {
return fmt.Errorf("读取 S0+S1+S2 失败: %w", err)
}
// 验证服务器版本
serverVersion := s0s1s2[0]
if serverVersion != RTMP_VERSION {
log.Printf("⚠️ [RTMP Handshake] 服务器版本: %d", serverVersion)
}
s1 := s0s1s2[1 : 1+HANDSHAKE_SIZE]
// 3. 发送 C2 (回显 S1)
c2 := make([]byte, HANDSHAKE_SIZE)
copy(c2[0:4], s1[0:4]) // 回显 S1 时间戳
binary.BigEndian.PutUint32(c2[4:8], binary.BigEndian.Uint32(c1[0:4])) // C1 时间戳
copy(c2[8:], s1[8:]) // 回显 S1 随机数据
if _, err := conn.Write(c2); err != nil {
return fmt.Errorf("发送 C2 失败: %w", err)
}
return nil
}

View File

@@ -0,0 +1,509 @@
/**
* package mediaserver
*
* 自实现的 RTMP 协议处理
* 用于替代 github.com/yutopp/go-rtmp解决 SetChunkSize 导致的 panic 问题
*/
package mediaserver
import (
"bufio"
"encoding/binary"
"fmt"
"io"
"log"
"net"
"sync"
)
// RTMP 消息类型常量
const (
RTMP_MSG_CHUNK_SIZE = 1 // SetChunkSize
RTMP_MSG_ABORT = 2 // Abort Message
RTMP_MSG_ACK = 3 // Acknowledgement
RTMP_MSG_USER_CONTROL = 4 // User Control Message
RTMP_MSG_WIN_ACK_SIZE = 5 // Window Acknowledgement Size
RTMP_MSG_SET_PEER_BW = 6 // Set Peer Bandwidth
RTMP_MSG_AUDIO = 8 // Audio Message
RTMP_MSG_VIDEO = 9 // Video Message
RTMP_MSG_AMF3_DATA = 15 // AMF3 Data Message
RTMP_MSG_AMF3_SHARED_OBJ = 16 // AMF3 Shared Object Message
RTMP_MSG_AMF3_CMD = 17 // AMF3 Command Message
RTMP_MSG_AMF0_DATA = 18 // AMF0 Data Message (@setDataFrame)
RTMP_MSG_AMF0_SHARED_OBJ = 19 // AMF0 Shared Object Message
RTMP_MSG_AMF0_CMD = 20 // AMF0 Command Message (connect, publish, etc)
RTMP_MSG_AGGREGATE = 22 // Aggregate Message
)
// Chunk 格式类型
const (
CHUNK_FMT_0 = 0 // 11 bytes header
CHUNK_FMT_1 = 1 // 7 bytes header
CHUNK_FMT_2 = 2 // 3 bytes header
CHUNK_FMT_3 = 3 // 0 bytes header
)
// 默认值
const (
DEFAULT_CHUNK_SIZE = 128
MAX_CHUNK_SIZE = 65536
DEFAULT_WINDOW_SIZE = 2500000
RTMP_PROTOCOL_VERSION = 3
)
// RTMPMessage 表示一个完整的 RTMP 消息
type RTMPMessage struct {
ChunkStreamID uint32
Timestamp uint32
TypeID uint8
StreamID uint32
Data []byte
}
// ChunkHeader 表示 chunk 的头部信息
type ChunkHeader struct {
Format uint8 // 0-3
ChunkStreamID uint32 // 2-65599
Timestamp uint32 // 24-bit or 32-bit (extended)
MessageLength uint32 // 24-bit
MessageTypeID uint8
MessageSID uint32 // 32-bit little-endian
ExtendedTS bool // Timestamp >= 0xFFFFFF
}
// ChunkReader RTMP Chunk 读取器
type ChunkReader struct {
conn net.Conn
reader *bufio.Reader
chunkSize uint32
mu sync.RWMutex
// 缓存每个 chunk stream 的头部信息(用于 fmt 1/2/3
prevHeaders map[uint32]*ChunkHeader
// 缓存不完整的消息数据
messageBuffer map[uint32]*messageState
}
// messageState 追踪消息的读取状态
type messageState struct {
header *ChunkHeader
data []byte
bytesRead uint32
}
// ChunkWriter RTMP Chunk 写入器
type ChunkWriter struct {
conn net.Conn
writer *bufio.Writer
chunkSize uint32
mu sync.Mutex
}
// NewChunkReader 创建新的 Chunk 读取器
func NewChunkReader(conn net.Conn) *ChunkReader {
return &ChunkReader{
conn: conn,
reader: bufio.NewReaderSize(conn, 4096),
chunkSize: DEFAULT_CHUNK_SIZE,
prevHeaders: make(map[uint32]*ChunkHeader),
messageBuffer: make(map[uint32]*messageState),
}
}
// NewChunkWriter 创建新的 Chunk 写入器
func NewChunkWriter(conn net.Conn) *ChunkWriter {
return &ChunkWriter{
conn: conn,
writer: bufio.NewWriterSize(conn, 4096),
chunkSize: DEFAULT_CHUNK_SIZE,
}
}
// SetChunkSize 设置读取的 chunk 大小
func (r *ChunkReader) SetChunkSize(size uint32) {
r.mu.Lock()
defer r.mu.Unlock()
if size > 0 && size <= MAX_CHUNK_SIZE {
log.Printf("📝 [RTMP Protocol] ChunkReader: 更新 chunkSize %d -> %d", r.chunkSize, size)
r.chunkSize = size
}
}
// GetChunkSize 获取当前 chunk 大小
func (r *ChunkReader) GetChunkSize() uint32 {
r.mu.RLock()
defer r.mu.RUnlock()
return r.chunkSize
}
// SetChunkSize 设置写入的 chunk 大小
func (w *ChunkWriter) SetChunkSize(size uint32) {
w.mu.Lock()
defer w.mu.Unlock()
if size > 0 && size <= MAX_CHUNK_SIZE {
log.Printf("📝 [RTMP Protocol] ChunkWriter: 更新 chunkSize %d -> %d", w.chunkSize, size)
w.chunkSize = size
}
}
// ReadMessage 读取一个完整的 RTMP 消息
func (r *ChunkReader) ReadMessage() (*RTMPMessage, error) {
for {
// 读取 chunk header
header, err := r.readChunkHeader()
if err != nil {
return nil, err
}
// 获取或创建消息状态
// 关键修复:检查消息是否已完成,如果是则为新消息创建新状态
state, exists := r.messageBuffer[header.ChunkStreamID]
if !exists || state.bytesRead >= state.header.MessageLength {
// 新消息或上一条消息已完成,创建新状态
state = &messageState{
header: header,
data: make([]byte, 0, header.MessageLength),
bytesRead: 0,
}
r.messageBuffer[header.ChunkStreamID] = state
}
// 计算本次要读取的字节数
// 关键修复:使用 state.header.MessageLength 而非 header.MessageLength
// 因为后续 chunk (fmt 1/2/3) 的 header 可能从前一个 header 继承值
remaining := state.header.MessageLength - state.bytesRead
r.mu.RLock()
toRead := r.chunkSize
r.mu.RUnlock()
if remaining < toRead {
toRead = remaining
}
// toRead=0 说明 MessageLength=0这是无效的 RTMP 消息
// 直接返回错误而不是继续,避免失去同步
if toRead == 0 {
return nil, fmt.Errorf("invalid RTMP: toRead=0 on csid=%d msgLen=%d bytesRead=%d",
header.ChunkStreamID, state.header.MessageLength, state.bytesRead)
}
// 读取 chunk 数据
chunkData := make([]byte, toRead)
if _, err := io.ReadFull(r.reader, chunkData); err != nil {
return nil, fmt.Errorf("读取 chunk 数据失败: %w", err)
}
state.data = append(state.data, chunkData...)
state.bytesRead += toRead
// 检查消息是否完整
if state.bytesRead >= state.header.MessageLength {
msg := &RTMPMessage{
ChunkStreamID: header.ChunkStreamID,
Timestamp: state.header.Timestamp,
TypeID: state.header.MessageTypeID,
StreamID: state.header.MessageSID,
Data: state.data,
}
// 清除消息缓冲
delete(r.messageBuffer, header.ChunkStreamID)
return msg, nil
}
}
}
// readChunkHeader 读取 chunk 头部
func (r *ChunkReader) readChunkHeader() (*ChunkHeader, error) {
// 读取第一个字节Basic Header (1-3 bytes)
firstByte, err := r.reader.ReadByte()
if err != nil {
return nil, err
}
format := (firstByte >> 6) & 0x03
csid := uint32(firstByte & 0x3F)
// 扩展 chunk stream ID
if csid == 0 {
// 2 byte header
secondByte, err := r.reader.ReadByte()
if err != nil {
return nil, err
}
csid = uint32(secondByte) + 64
} else if csid == 1 {
// 3 byte header
bytes := make([]byte, 2)
if _, err := io.ReadFull(r.reader, bytes); err != nil {
return nil, err
}
csid = uint32(bytes[0]) + uint32(bytes[1])*256 + 64
}
// 获取上一个头部(用于 fmt 1/2/3
prevHeader := r.prevHeaders[csid]
if prevHeader == nil {
prevHeader = &ChunkHeader{
ChunkStreamID: csid,
}
}
header := &ChunkHeader{
Format: format,
ChunkStreamID: csid,
Timestamp: prevHeader.Timestamp,
MessageLength: prevHeader.MessageLength,
MessageTypeID: prevHeader.MessageTypeID,
MessageSID: prevHeader.MessageSID,
}
// 根据 format 读取 Message Header
switch format {
case CHUNK_FMT_0:
// 11 bytes: timestamp(3) + length(3) + typeID(1) + streamID(4)
data := make([]byte, 11)
if _, err := io.ReadFull(r.reader, data); err != nil {
return nil, fmt.Errorf("读取 fmt0 header 失败: %w", err)
}
header.Timestamp = uint32(data[0])<<16 | uint32(data[1])<<8 | uint32(data[2])
header.MessageLength = uint32(data[3])<<16 | uint32(data[4])<<8 | uint32(data[5])
header.MessageTypeID = data[6]
header.MessageSID = binary.LittleEndian.Uint32(data[7:11])
case CHUNK_FMT_1:
// 7 bytes: timestamp delta(3) + length(3) + typeID(1)
data := make([]byte, 7)
if _, err := io.ReadFull(r.reader, data); err != nil {
return nil, fmt.Errorf("读取 fmt1 header 失败: %w", err)
}
delta := uint32(data[0])<<16 | uint32(data[1])<<8 | uint32(data[2])
header.Timestamp = prevHeader.Timestamp + delta
header.MessageLength = uint32(data[3])<<16 | uint32(data[4])<<8 | uint32(data[5])
header.MessageTypeID = data[6]
case CHUNK_FMT_2:
// 验证fmt 2 必须有有效的 prevHeaderMessageLength 从 prevHeader 继承)
if prevHeader.MessageLength == 0 {
return nil, fmt.Errorf("invalid RTMP: fmt 2 on csid %d without valid prevHeader (msgLen=0)", csid)
}
// 3 bytes: timestamp delta(3)
data := make([]byte, 3)
if _, err := io.ReadFull(r.reader, data); err != nil {
return nil, fmt.Errorf("读取 fmt2 header 失败: %w", err)
}
delta := uint32(data[0])<<16 | uint32(data[1])<<8 | uint32(data[2])
header.Timestamp = prevHeader.Timestamp + delta
case CHUNK_FMT_3:
// 验证fmt 3 必须有有效的 prevHeader所有字段从 prevHeader 继承)
if prevHeader.MessageLength == 0 {
return nil, fmt.Errorf("invalid RTMP: fmt 3 on csid %d without valid prevHeader (msgLen=0)", csid)
}
// 0 bytes: 使用上一个头部的所有字段
// 已经复制了 prevHeader 的值
}
// 检查是否有扩展时间戳
if header.Timestamp == 0xFFFFFF {
header.ExtendedTS = true
extTS := make([]byte, 4)
if _, err := io.ReadFull(r.reader, extTS); err != nil {
return nil, fmt.Errorf("读取扩展时间戳失败: %w", err)
}
header.Timestamp = binary.BigEndian.Uint32(extTS)
}
// 保存当前头部供后续 chunk 使用
r.prevHeaders[csid] = header
return header, nil
}
// WriteMessage 写入一个完整的 RTMP 消息
func (w *ChunkWriter) WriteMessage(msg *RTMPMessage) error {
w.mu.Lock()
defer w.mu.Unlock()
data := msg.Data
dataLen := uint32(len(data))
offset := uint32(0)
firstChunk := true
for offset < dataLen {
// 计算本次写入的字节数
remaining := dataLen - offset
toWrite := w.chunkSize
if remaining < toWrite {
toWrite = remaining
}
// 写入 chunk header
if firstChunk {
// fmt 0: 完整头部
if err := w.writeChunkHeader0(msg, dataLen); err != nil {
return err
}
firstChunk = false
} else {
// fmt 3: 无头部
if err := w.writeChunkHeader3(msg.ChunkStreamID); err != nil {
return err
}
}
// 写入数据
if _, err := w.writer.Write(data[offset : offset+toWrite]); err != nil {
return err
}
offset += toWrite
}
return w.writer.Flush()
}
// writeChunkHeader0 写入 fmt 0 头部
func (w *ChunkWriter) writeChunkHeader0(msg *RTMPMessage, dataLen uint32) error {
// Basic Header
csid := msg.ChunkStreamID
if csid < 64 {
if err := w.writer.WriteByte(byte(csid)); err != nil {
return err
}
} else if csid < 320 {
if _, err := w.writer.Write([]byte{0, byte(csid - 64)}); err != nil {
return err
}
} else {
csid -= 64
if _, err := w.writer.Write([]byte{1, byte(csid & 0xFF), byte(csid >> 8)}); err != nil {
return err
}
}
// Message Header (11 bytes)
header := make([]byte, 11)
ts := msg.Timestamp
if ts >= 0xFFFFFF {
ts = 0xFFFFFF
}
header[0] = byte(ts >> 16)
header[1] = byte(ts >> 8)
header[2] = byte(ts)
header[3] = byte(dataLen >> 16)
header[4] = byte(dataLen >> 8)
header[5] = byte(dataLen)
header[6] = msg.TypeID
binary.LittleEndian.PutUint32(header[7:11], msg.StreamID)
if _, err := w.writer.Write(header); err != nil {
return err
}
// Extended Timestamp
if msg.Timestamp >= 0xFFFFFF {
extTS := make([]byte, 4)
binary.BigEndian.PutUint32(extTS, msg.Timestamp)
if _, err := w.writer.Write(extTS); err != nil {
return err
}
}
return nil
}
// writeChunkHeader3 写入 fmt 3 头部
func (w *ChunkWriter) writeChunkHeader3(csid uint32) error {
// Basic Header with fmt = 3
if csid < 64 {
return w.writer.WriteByte(byte(0xC0 | csid))
} else if csid < 320 {
_, err := w.writer.Write([]byte{0xC0, byte(csid - 64)})
return err
} else {
csid -= 64
_, err := w.writer.Write([]byte{0xC1, byte(csid & 0xFF), byte(csid >> 8)})
return err
}
}
// WriteSetChunkSize 发送 SetChunkSize 消息
func (w *ChunkWriter) WriteSetChunkSize(size uint32) error {
data := make([]byte, 4)
binary.BigEndian.PutUint32(data, size)
msg := &RTMPMessage{
ChunkStreamID: 2,
Timestamp: 0,
TypeID: RTMP_MSG_CHUNK_SIZE,
StreamID: 0,
Data: data,
}
if err := w.WriteMessage(msg); err != nil {
return err
}
w.SetChunkSize(size)
return nil
}
// WriteWindowAckSize 发送 Window Acknowledgement Size 消息
func (w *ChunkWriter) WriteWindowAckSize(size uint32) error {
data := make([]byte, 4)
binary.BigEndian.PutUint32(data, size)
return w.WriteMessage(&RTMPMessage{
ChunkStreamID: 2,
Timestamp: 0,
TypeID: RTMP_MSG_WIN_ACK_SIZE,
StreamID: 0,
Data: data,
})
}
// WriteSetPeerBandwidth 发送 Set Peer Bandwidth 消息
func (w *ChunkWriter) WriteSetPeerBandwidth(size uint32, limitType uint8) error {
data := make([]byte, 5)
binary.BigEndian.PutUint32(data, size)
data[4] = limitType
return w.WriteMessage(&RTMPMessage{
ChunkStreamID: 2,
Timestamp: 0,
TypeID: RTMP_MSG_SET_PEER_BW,
StreamID: 0,
Data: data,
})
}
// WriteUserControl 发送 User Control Message
func (w *ChunkWriter) WriteUserControl(eventType uint16, data []byte) error {
payload := make([]byte, 2+len(data))
binary.BigEndian.PutUint16(payload, eventType)
copy(payload[2:], data)
return w.WriteMessage(&RTMPMessage{
ChunkStreamID: 2,
Timestamp: 0,
TypeID: RTMP_MSG_USER_CONTROL,
StreamID: 0,
Data: payload,
})
}
// WriteStreamBegin 发送 Stream Begin 事件
func (w *ChunkWriter) WriteStreamBegin(streamID uint32) error {
data := make([]byte, 4)
binary.BigEndian.PutUint32(data, streamID)
return w.WriteUserControl(0, data) // 0 = StreamBegin
}
// WriteCommand 发送 AMF0 命令消息
func (w *ChunkWriter) WriteCommand(csid uint32, streamID uint32, data []byte) error {
return w.WriteMessage(&RTMPMessage{
ChunkStreamID: csid,
Timestamp: 0,
TypeID: RTMP_MSG_AMF0_CMD,
StreamID: streamID,
Data: data,
})
}

View File

@@ -18,11 +18,12 @@ import (
// MediaServer 媒体服务器主结构
type MediaServer struct {
sfu *SFUServer // WebRTC SFU 服务
rtmp *RTMPServer // RTMP 服务
rooms map[string]*Room // 通话房间管
mu sync.RWMutex // 房间锁
config *MediaServerConfig // 配置
sfu *SFUServer // WebRTC SFU 服务
rtmp *RTMPServer // RTMP 服务
wsProxy *WebMToRTMPProxy // WebSocket to RTMP 代
rooms map[string]*Room // 通话房间管理
mu sync.RWMutex // 房间锁
config *MediaServerConfig // 配置
}
// MediaServerConfig 媒体服务器配置
@@ -65,6 +66,8 @@ func newMediaServer() *MediaServer {
if config.Enabled {
ms.sfu = NewSFUServer(config)
ms.rtmp = NewRTMPServer(config)
// 创建 WebSocket to RTMP 代理
ms.wsProxy = NewWebMToRTMPProxy(ms.rtmp)
}
return ms
@@ -208,6 +211,11 @@ func (ms *MediaServer) GetRTMP() *RTMPServer {
return ms.rtmp
}
// GetWSProxy 获取 WebSocket to RTMP 代理
func (ms *MediaServer) GetWSProxy() *WebMToRTMPProxy {
return ms.wsProxy
}
// GetRoomCount 获取房间数量
func (ms *MediaServer) GetRoomCount() int {
ms.mu.RLock()

View File

@@ -0,0 +1,736 @@
/**
* package mediaserver
*
* WebSocket to RTMP 代理服务
* 接收 Web端 通过 WebSocket 发送的 WebM 数据,转换为 FLV 并推送到 RTMP
*
* 技术栈:
* - WebSocket: gorilla/websocket
* - WebM 解析: github.com/at-wat/ebml-go
* - FLV 封装: github.com/yutopp/go-flv
* - RTMP 推送: 通过内部流管道
*/
package mediaserver
import (
"bytes"
"encoding/binary"
"fmt"
"io"
"log"
"net/http"
"sync"
"time"
"github.com/at-wat/ebml-go"
"github.com/gorilla/websocket"
)
// WebMToRTMPProxy WebSocket to RTMP 代理
type WebMToRTMPProxy struct {
rtmpServer *RTMPServer
upgrader websocket.Upgrader
sessions map[string]*ProxySession
sessionsMu sync.RWMutex
}
// ProxySession 代理会话
type ProxySession struct {
ID string
StreamID string
UserID string
RoomID string
Conn *websocket.Conn
Stream *RTMPStream
StopChan chan struct{}
StartTime time.Time
// WebM 解析状态
webmBuffer *bytes.Buffer
headerParsed bool
videoTrackNum uint64
audioTrackNum uint64
// 时间戳
baseTimestamp uint32
lastTimestamp uint32
}
// WebMHeader WebM 文件头信息
type WebMHeader struct {
EBMLVersion uint64 `ebml:"EBMLVersion"`
EBMLReadVersion uint64 `ebml:"EBMLReadVersion"`
EBMLMaxIDLength uint64 `ebml:"EBMLMaxIDLength"`
EBMLMaxSizeLength uint64 `ebml:"EBMLMaxSizeLength"`
DocType string `ebml:"DocType"`
DocTypeVersion uint64 `ebml:"DocTypeVersion"`
DocTypeReadVersion uint64 `ebml:"DocTypeReadVersion"`
}
// WebMSegment WebM Segment
type WebMSegment struct {
Info WebMSegmentInfo `ebml:"Info"`
Tracks WebMTracks `ebml:"Tracks"`
Cluster []WebMCluster `ebml:"Cluster"`
}
// WebMSegmentInfo Segment 信息
type WebMSegmentInfo struct {
TimecodeScale uint64 `ebml:"TimecodeScale"`
Duration float64 `ebml:"Duration,omitempty"`
MuxingApp string `ebml:"MuxingApp,omitempty"`
WritingApp string `ebml:"WritingApp,omitempty"`
}
// WebMTracks 轨道信息
type WebMTracks struct {
TrackEntry []WebMTrackEntry `ebml:"TrackEntry"`
}
// WebMTrackEntry 轨道条目
type WebMTrackEntry struct {
TrackNumber uint64 `ebml:"TrackNumber"`
TrackType uint64 `ebml:"TrackType"` // 1=video, 2=audio
CodecID string `ebml:"CodecID"`
Video *WebMVideoTrack `ebml:"Video,omitempty"`
Audio *WebMAudioTrack `ebml:"Audio,omitempty"`
}
// WebMVideoTrack 视频轨道
type WebMVideoTrack struct {
PixelWidth uint64 `ebml:"PixelWidth"`
PixelHeight uint64 `ebml:"PixelHeight"`
}
// WebMAudioTrack 音频轨道
type WebMAudioTrack struct {
SamplingFrequency float64 `ebml:"SamplingFrequency"`
Channels uint64 `ebml:"Channels"`
BitDepth uint64 `ebml:"BitDepth,omitempty"`
}
// WebMCluster WebM Cluster
type WebMCluster struct {
Timecode uint64 `ebml:"Timecode"`
SimpleBlock []ebml.Block `ebml:"SimpleBlock,omitempty"`
}
// NewWebMToRTMPProxy 创建代理服务
func NewWebMToRTMPProxy(rtmpServer *RTMPServer) *WebMToRTMPProxy {
return &WebMToRTMPProxy{
rtmpServer: rtmpServer,
upgrader: websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool {
return true // 允许跨域
},
ReadBufferSize: 1024 * 1024,
WriteBufferSize: 1024 * 1024,
},
sessions: make(map[string]*ProxySession),
}
}
// HandleWebSocket 处理 WebSocket 连接
func (p *WebMToRTMPProxy) HandleWebSocket(w http.ResponseWriter, r *http.Request) {
log.Printf("🔌 [WSProxy] 收到 WebSocket 连接请求: %s from %s", r.URL.String(), r.RemoteAddr)
// 检查 WebSocket 升级头
if r.Header.Get("Upgrade") != "websocket" {
log.Printf("❌ [WSProxy] 不是 WebSocket 请求: Upgrade=%s", r.Header.Get("Upgrade"))
http.Error(w, "Not a WebSocket request", http.StatusBadRequest)
return
}
// 获取参数
streamID := r.URL.Query().Get("stream_id")
userID := r.URL.Query().Get("user_id")
roomID := r.URL.Query().Get("room_id")
token := r.URL.Query().Get("token")
log.Printf("📋 [WSProxy] 参数: stream_id=%s user_id=%s room_id=%s token=%v",
streamID, userID, roomID, token != "")
if streamID == "" || userID == "" {
log.Printf("❌ [WSProxy] 缺少必要参数")
http.Error(w, "Missing stream_id or user_id", http.StatusBadRequest)
return
}
// 验证 token简化版生产环境需要更严格的验证
if token == "" {
log.Printf("⚠️ [WSProxy] 缺少 token: stream=%s", streamID)
}
// 升级到 WebSocket
log.Printf("🔄 [WSProxy] 尝试 WebSocket 升级...")
conn, err := p.upgrader.Upgrade(w, r, nil)
if err != nil {
log.Printf("❌ [WSProxy] WebSocket 升级失败: %v (可能原因: 中间件干扰、响应已写入)", err)
return
}
log.Printf("✅ [WSProxy] WebSocket 升级成功 | Stream:%s User:%s Room:%s", streamID, userID, roomID)
// 获取或创建 RTMP 流
stream := p.rtmpServer.GetStream(streamID)
if stream == nil {
// 自动创建流
var err error
stream, err = p.rtmpServer.GenerateStreamURLs(roomID, userID)
if err != nil {
log.Printf("❌ [WSProxy] 创建流失败: %v", err)
conn.Close()
return
}
// 更新 streamID 为实际生成的
streamID = stream.ID
}
// 创建会话
session := &ProxySession{
ID: fmt.Sprintf("%s_%d", userID, time.Now().UnixNano()),
StreamID: streamID,
UserID: userID,
RoomID: roomID,
Conn: conn,
Stream: stream,
StopChan: make(chan struct{}),
StartTime: time.Now(),
webmBuffer: bytes.NewBuffer(nil),
headerParsed: false,
}
// 注册会话
p.sessionsMu.Lock()
p.sessions[session.ID] = session
p.sessionsMu.Unlock()
// 激活流
stream.mu.Lock()
stream.IsActive = true
stream.mu.Unlock()
// 处理连接
go p.handleSession(session)
}
// handleSession 处理会话
func (p *WebMToRTMPProxy) handleSession(session *ProxySession) {
log.Printf("▶️ [WSProxy] 开始处理会话 | Stream:%s User:%s", session.StreamID, session.UserID)
messageCount := 0
totalBytes := int64(0)
defer func() {
// 清理
p.sessionsMu.Lock()
delete(p.sessions, session.ID)
p.sessionsMu.Unlock()
session.Conn.Close()
close(session.StopChan)
// 标记流为非活动
if session.Stream != nil {
session.Stream.mu.Lock()
session.Stream.IsActive = false
session.Stream.mu.Unlock()
}
log.Printf("🔌 [WSProxy] 连接关闭 | Stream:%s User:%s | 收到消息:%d 总字节:%d",
session.StreamID, session.UserID, messageCount, totalBytes)
}()
for {
select {
case <-session.StopChan:
log.Printf("⏹️ [WSProxy] 收到停止信号 | Stream:%s", session.StreamID)
return
default:
}
// 读取 WebSocket 消息
messageType, data, err := session.Conn.ReadMessage()
if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) {
log.Printf("⚠️ [WSProxy] 读取错误: %v | Stream:%s", err, session.StreamID)
} else {
log.Printf(" [WSProxy] 连接关闭: %v | Stream:%s", err, session.StreamID)
}
return
}
messageCount++
totalBytes += int64(len(data))
// 记录首次收到消息
if messageCount == 1 {
log.Printf("📥 [WSProxy] 首次收到消息 | Stream:%s Type:%d Size:%d", session.StreamID, messageType, len(data))
}
// 每 100 条消息记录一次统计
if messageCount%100 == 0 {
log.Printf("📊 [WSProxy] 消息统计 | Stream:%s Count:%d TotalBytes:%d", session.StreamID, messageCount, totalBytes)
}
if messageType != websocket.BinaryMessage {
log.Printf("⚠️ [WSProxy] 忽略非二进制消息 | Type:%d", messageType)
continue
}
// 处理 WebM 数据
if err := p.processWebMData(session, data); err != nil {
log.Printf("⚠️ [WSProxy] 处理 WebM 数据失败: %v", err)
}
}
}
// processWebMData 处理 WebM 数据
func (p *WebMToRTMPProxy) processWebMData(session *ProxySession, data []byte) error {
// 将数据追加到缓冲区
session.webmBuffer.Write(data)
// 尝试解析 WebM 数据
return p.parseAndConvert(session)
}
// parseAndConvert 解析 WebM 并转换为 FLV
func (p *WebMToRTMPProxy) parseAndConvert(session *ProxySession) error {
bufData := session.webmBuffer.Bytes()
if len(bufData) < 4 {
return nil // 数据不足
}
// 检查 EBML 头部
if !session.headerParsed {
// 尝试解析头部
if err := p.parseWebMHeader(session, bufData); err != nil {
// 头部不完整,等待更多数据
return nil
}
session.headerParsed = true
log.Printf("📦 [WSProxy] WebM 头部解析完成 | Stream:%s", session.StreamID)
}
// 解析 Cluster 并转换
return p.parseWebMClusters(session)
}
// parseWebMHeader 解析 WebM 头部
func (p *WebMToRTMPProxy) parseWebMHeader(session *ProxySession, data []byte) error {
reader := bytes.NewReader(data)
// 解析 EBML 头
var header struct {
EBML struct {
EBMLVersion uint64 `ebml:"EBMLVersion"`
DocType string `ebml:"DocType"`
} `ebml:"EBML"`
}
if err := ebml.Unmarshal(reader, &header); err != nil {
return err
}
log.Printf("📦 [WSProxy] WebM DocType: %s", header.EBML.DocType)
return nil
}
// parseWebMClusters 解析 WebM Clusters
func (p *WebMToRTMPProxy) parseWebMClusters(session *ProxySession) error {
// 简化处理:直接将 WebM 数据转换为 FLV
// 实际实现需要完整解析 WebM 的 Cluster/SimpleBlock
bufData := session.webmBuffer.Bytes()
if len(bufData) < 100 {
return nil
}
// 寻找 Cluster 标记 (0x1F43B675)
clusterMarker := []byte{0x1F, 0x43, 0xB6, 0x75}
for {
idx := bytes.Index(bufData, clusterMarker)
if idx == -1 || idx+12 > len(bufData) {
break
}
// 解析 Cluster 大小
sizeStart := idx + 4
clusterSize, bytesRead := readVarInt(bufData[sizeStart:])
if bytesRead == 0 {
break
}
totalSize := idx + 4 + bytesRead + int(clusterSize)
if totalSize > len(bufData) {
// Cluster 不完整
break
}
// 提取 Cluster 数据
clusterData := bufData[idx:totalSize]
// 转换为 FLV 并广播
if err := p.convertClusterToFLV(session, clusterData); err != nil {
log.Printf("⚠️ [WSProxy] 转换 Cluster 失败: %v", err)
}
// 从缓冲区移除已处理的数据
bufData = bufData[totalSize:]
}
// 更新缓冲区
session.webmBuffer.Reset()
session.webmBuffer.Write(bufData)
return nil
}
// readVarInt 读取 EBML 变长整数
func readVarInt(data []byte) (uint64, int) {
if len(data) == 0 {
return 0, 0
}
first := data[0]
var length int
var mask byte
switch {
case first&0x80 != 0:
length = 1
mask = 0x7F
case first&0x40 != 0:
length = 2
mask = 0x3F
case first&0x20 != 0:
length = 3
mask = 0x1F
case first&0x10 != 0:
length = 4
mask = 0x0F
case first&0x08 != 0:
length = 5
mask = 0x07
case first&0x04 != 0:
length = 6
mask = 0x03
case first&0x02 != 0:
length = 7
mask = 0x01
case first&0x01 != 0:
length = 8
mask = 0x00
default:
return 0, 0
}
if len(data) < length {
return 0, 0
}
value := uint64(data[0] & mask)
for i := 1; i < length; i++ {
value = (value << 8) | uint64(data[i])
}
return value, length
}
// convertClusterToFLV 将 WebM Cluster 转换为 FLV 并广播
func (p *WebMToRTMPProxy) convertClusterToFLV(session *ProxySession, clusterData []byte) error {
// 解析 Cluster 时间戳
timecodeMarker := []byte{0xE7} // Timecode element ID
timecodeIdx := bytes.Index(clusterData, timecodeMarker)
var timestamp uint32 = session.lastTimestamp
if timecodeIdx != -1 && timecodeIdx+1 < len(clusterData) {
tcSize, bytesRead := readVarInt(clusterData[timecodeIdx+1:])
if bytesRead > 0 && timecodeIdx+1+bytesRead+int(tcSize) <= len(clusterData) {
tcData := clusterData[timecodeIdx+1+bytesRead : timecodeIdx+1+bytesRead+int(tcSize)]
for _, b := range tcData {
timestamp = (timestamp << 8) | uint32(b)
}
}
}
// 解析 SimpleBlock
simpleBlockMarker := []byte{0xA3} // SimpleBlock element ID
blockData := clusterData
for {
blockIdx := bytes.Index(blockData, simpleBlockMarker)
if blockIdx == -1 || blockIdx+1 >= len(blockData) {
break
}
// 解析 block 大小
sizeStart := blockIdx + 1
blockSize, bytesRead := readVarInt(blockData[sizeStart:])
if bytesRead == 0 {
break
}
dataStart := sizeStart + bytesRead
dataEnd := dataStart + int(blockSize)
if dataEnd > len(blockData) {
break
}
// 提取 block 内容
block := blockData[dataStart:dataEnd]
if len(block) < 4 {
blockData = blockData[dataEnd:]
continue
}
// 解析 track number变长
trackNum, trackBytes := readVarInt(block)
if trackBytes == 0 {
blockData = blockData[dataEnd:]
continue
}
// 解析相对时间戳2 bytes, big-endian
if trackBytes+2 > len(block) {
blockData = blockData[dataEnd:]
continue
}
relativeTimestamp := binary.BigEndian.Uint16(block[trackBytes : trackBytes+2])
// 解析 flags (1 byte)
flagsIdx := trackBytes + 2
if flagsIdx >= len(block) {
blockData = blockData[dataEnd:]
continue
}
flags := block[flagsIdx]
// 帧数据
frameData := block[flagsIdx+1:]
if len(frameData) == 0 {
blockData = blockData[dataEnd:]
continue
}
// 计算绝对时间戳
absoluteTimestamp := timestamp + uint32(relativeTimestamp)
session.lastTimestamp = absoluteTimestamp
// 判断是视频还是音频(简化:假设 track 1 是视频track 2 是音频)
isVideo := trackNum == 1
isKeyframe := (flags & 0x80) != 0
// 创建 FLV tag 并广播
var flvTag []byte
if isVideo {
flvTag = p.createFLVVideoTag(absoluteTimestamp, frameData, isKeyframe)
} else {
flvTag = p.createFLVAudioTag(absoluteTimestamp, frameData)
}
if flvTag != nil {
p.broadcastFLVTag(session, flvTag)
}
blockData = blockData[dataEnd:]
}
return nil
}
// createFLVVideoTag 创建 FLV 视频 tag
// VP8 -> FLV (需要转换为 H.264,这里简化处理)
func (p *WebMToRTMPProxy) createFLVVideoTag(timestamp uint32, data []byte, isKeyframe bool) []byte {
// 注意VP8 不能直接封装到 FLV需要转码为 H.264
// 这里使用一个简化的方案:将 VP8 数据作为自定义格式封装
// 实际生产环境需要使用 FFmpeg 或硬件编码器进行转码
// FLV Video Tag Header:
// FrameType (4 bits): 1=keyframe, 2=inter frame
// CodecID (4 bits): 7=AVC (H.264)
// 由于 VP8 无法直接放入 FLV这里使用一个变通方案
// 将 VP8 数据标记为私有编码格式
frameType := byte(2) // inter frame
if isKeyframe {
frameType = 1 // keyframe
}
// 使用 CodecID=12 (VP8 - 非标准,仅用于内部传输)
// 或者可以考虑在服务端进行实时转码
codecID := byte(12) // 自定义VP8
header := (frameType << 4) | codecID
// 构建完整数据
videoData := make([]byte, 1+len(data))
videoData[0] = header
copy(videoData[1:], data)
return p.createFLVTag(9, timestamp, videoData) // 9 = video
}
// createFLVAudioTag 创建 FLV 音频 tag
// Opus -> FLV (需要转换为 AAC这里简化处理)
func (p *WebMToRTMPProxy) createFLVAudioTag(timestamp uint32, data []byte) []byte {
// 注意Opus 不能直接封装到 FLV需要转码为 AAC
// 这里使用简化方案
// FLV Audio Tag Header:
// SoundFormat (4 bits): 10=AAC, 13=Opus (非标准)
// SoundRate (2 bits): 3=44kHz
// SoundSize (1 bit): 1=16-bit
// SoundType (1 bit): 1=stereo
// 使用自定义格式标记 Opus
soundFormat := byte(13) // 自定义Opus
soundRate := byte(3) // 44kHz
soundSize := byte(1) // 16-bit
soundType := byte(1) // stereo
header := (soundFormat << 4) | (soundRate << 2) | (soundSize << 1) | soundType
// 构建完整数据
audioData := make([]byte, 1+len(data))
audioData[0] = header
copy(audioData[1:], data)
return p.createFLVTag(8, timestamp, audioData) // 8 = audio
}
// createFLVTag 创建 FLV tag
func (p *WebMToRTMPProxy) createFLVTag(tagType byte, timestamp uint32, data []byte) []byte {
dataSize := len(data)
tagSize := 11 + dataSize + 4
tag := make([]byte, tagSize)
// Tag type
tag[0] = tagType
// Data size (24-bit big-endian)
tag[1] = byte((dataSize >> 16) & 0xff)
tag[2] = byte((dataSize >> 8) & 0xff)
tag[3] = byte(dataSize & 0xff)
// Timestamp (24-bit big-endian)
tag[4] = byte((timestamp >> 16) & 0xff)
tag[5] = byte((timestamp >> 8) & 0xff)
tag[6] = byte(timestamp & 0xff)
// Timestamp extended
tag[7] = byte((timestamp >> 24) & 0xff)
// Stream ID (always 0)
tag[8] = 0
tag[9] = 0
tag[10] = 0
// Data
copy(tag[11:], data)
// Previous tag size
prevTagSize := 11 + dataSize
tag[11+dataSize] = byte((prevTagSize >> 24) & 0xff)
tag[11+dataSize+1] = byte((prevTagSize >> 16) & 0xff)
tag[11+dataSize+2] = byte((prevTagSize >> 8) & 0xff)
tag[11+dataSize+3] = byte(prevTagSize & 0xff)
return tag
}
// broadcastFLVTag 广播 FLV tag 到订阅者
func (p *WebMToRTMPProxy) broadcastFLVTag(session *ProxySession, tag []byte) {
if session.Stream == nil {
return
}
// 缓存到 GOP
session.Stream.mu.Lock()
session.Stream.gopCache = append(session.Stream.gopCache, tag)
if len(session.Stream.gopCache) > 300 {
session.Stream.gopCache = session.Stream.gopCache[len(session.Stream.gopCache)-300:]
}
session.Stream.mu.Unlock()
// 广播到 HTTP-FLV 订阅者
p.rtmpServer.subMu.RLock()
subs, exists := p.rtmpServer.subscribers[session.StreamID]
if exists {
for _, sub := range subs {
select {
case sub.DataChan <- tag:
default:
// 缓冲区满,跳过
}
}
}
p.rtmpServer.subMu.RUnlock()
}
// GetSession 获取会话
func (p *WebMToRTMPProxy) GetSession(sessionID string) *ProxySession {
p.sessionsMu.RLock()
defer p.sessionsMu.RUnlock()
return p.sessions[sessionID]
}
// GetSessionsByStream 获取流的所有会话
func (p *WebMToRTMPProxy) GetSessionsByStream(streamID string) []*ProxySession {
p.sessionsMu.RLock()
defer p.sessionsMu.RUnlock()
sessions := make([]*ProxySession, 0)
for _, s := range p.sessions {
if s.StreamID == streamID {
sessions = append(sessions, s)
}
}
return sessions
}
// CloseSession 关闭会话
func (p *WebMToRTMPProxy) CloseSession(sessionID string) {
p.sessionsMu.Lock()
session, exists := p.sessions[sessionID]
if exists {
delete(p.sessions, sessionID)
}
p.sessionsMu.Unlock()
if session != nil {
select {
case <-session.StopChan:
default:
close(session.StopChan)
}
session.Conn.Close()
}
}
// CloseAllSessions 关闭所有会话
func (p *WebMToRTMPProxy) CloseAllSessions() {
p.sessionsMu.Lock()
sessions := make([]*ProxySession, 0, len(p.sessions))
for _, s := range p.sessions {
sessions = append(sessions, s)
}
p.sessions = make(map[string]*ProxySession)
p.sessionsMu.Unlock()
for _, s := range sessions {
select {
case <-s.StopChan:
default:
close(s.StopChan)
}
s.Conn.Close()
}
}
// 确保导入被使用
var _ io.Reader = (*bytes.Reader)(nil)

View File

@@ -5,6 +5,7 @@
package middleware
import (
"log"
"net/http"
"xk-websocket-v2/internal/utils"
@@ -60,6 +61,9 @@ func JWTAuthMiddleware() gin.HandlerFunc {
return
}
// 调试日志打印解析出的用户ID帮助排查JWT问题
log.Printf("🔑 [JWT] 请求: %s %s | 解析出的用户ID: %s", c.Request.Method, c.Request.URL.Path, userID)
// 步骤6: 将解析出的用户ID注入到Context中
// 后续处理器可以通过 c.Get("user_id") 获取当前用户ID
c.Set("user_id", userID)

View File

@@ -40,6 +40,12 @@ var skipResponseBodyRoutes = []string{
"/static/",
}
// 需要完全跳过中间件的路由(如 WebSocket
var skipMiddlewareRoutes = []string{
"/ws",
"/api/call/ws-push",
}
// isBinaryContentType 检测是否为二进制 Content-Type
func isBinaryContentType(contentType string) bool {
contentType = strings.ToLower(contentType)
@@ -113,6 +119,16 @@ func NewRequestLogMiddleware(db *gorm.DB) *RequestLogMiddleware {
return &RequestLogMiddleware{DB: db}
}
// shouldSkipMiddleware 检测是否应该完全跳过中间件
func shouldSkipMiddleware(path string) bool {
for _, route := range skipMiddlewareRoutes {
if strings.HasPrefix(path, route) {
return true
}
}
return false
}
// Handler 中间件处理函数
func (m *RequestLogMiddleware) Handler() gin.HandlerFunc {
return func(c *gin.Context) {
@@ -122,6 +138,12 @@ func (m *RequestLogMiddleware) Handler() gin.HandlerFunc {
return
}
// 跳过 WebSocket 路由(包装 Writer 会干扰 WebSocket 升级)
if shouldSkipMiddleware(c.Request.URL.Path) {
c.Next()
return
}
// 获取请求IP
ip := getClientIP(c)