diff --git a/cmd/server/build.bat b/cmd/server/build.bat index a51a9ca..f871aa6 100644 --- a/cmd/server/build.bat +++ b/cmd/server/build.bat @@ -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. diff --git a/cmd/server/main.go b/cmd/server/main.go index 6905a12..4f277df 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -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") { diff --git a/go.mod b/go.mod index c447ce4..801a33d 100644 --- a/go.mod +++ b/go.mod @@ -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 diff --git a/go.sum b/go.sum index f2d6719..63fb6ad 100644 --- a/go.sum +++ b/go.sum @@ -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= diff --git a/internal/api/call_handler.go b/internal/api/call_handler.go index 0afa302..fee462e 100644 --- a/internal/api/call_handler.go +++ b/internal/api/call_handler.go @@ -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) +} diff --git a/internal/mediaserver/room.go b/internal/mediaserver/room.go index ca8cd5f..d3fc532 100644 --- a/internal/mediaserver/room.go +++ b/internal/mediaserver/room.go @@ -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 } diff --git a/internal/mediaserver/rtmp.go b/internal/mediaserver/rtmp.go index 85b541f..c0c928f 100644 --- a/internal/mediaserver/rtmp.go +++ b/internal/mediaserver/rtmp.go @@ -2,7 +2,7 @@ * package mediaserver * * RTMP 服务(用于微信小程序 live-pusher/live-player) - * 使用 github.com/yutopp/go-rtmp 实现完整的 RTMP 协议 + * 自实现 RTMP 协议,解决第三方库的 SetChunkSize panic 问题 * * 功能: * 1. 接收小程序推流(publish) @@ -12,12 +12,11 @@ package mediaserver import ( - "bytes" "crypto/hmac" "crypto/sha1" "encoding/base64" + "encoding/binary" "fmt" - "io" "log" "net" "net/http" @@ -26,10 +25,6 @@ import ( "time" "github.com/spf13/viper" - "github.com/yutopp/go-flv" - flvtag "github.com/yutopp/go-flv/tag" - "github.com/yutopp/go-rtmp" - rtmpmsg "github.com/yutopp/go-rtmp/message" ) // RTMPServer RTMP 服务器 @@ -43,7 +38,6 @@ type RTMPServer struct { // 服务器监听器 rtmpListener net.Listener httpServer *http.Server - rtmpServer *rtmp.Server // 订阅者管理 subscribers map[string]map[string]*Subscriber // streamID -> subscriberID -> subscriber @@ -126,59 +120,252 @@ func (r *RTMPServer) startRTMPServer() { log.Printf("🎬 [RTMP] RTMP 服务监听: %s", addr) - // 创建 RTMP 服务器 - r.rtmpServer = rtmp.NewServer(&rtmp.ServerConfig{ - OnConnect: func(conn net.Conn) (io.ReadWriteCloser, *rtmp.ConnConfig) { - log.Printf("🎬 [RTMP] 新连接: %s", conn.RemoteAddr().String()) - return conn, &rtmp.ConnConfig{ - Handler: &RTMPHandler{ - server: r, - conn: conn, - }, - ControlState: rtmp.StreamControlStateConfig{ - DefaultBandwidthWindowSize: 6 * 1024 * 1024, - }, + for { + conn, err := r.rtmpListener.Accept() + if err != nil { + r.mu.RLock() + running := r.running + r.mu.RUnlock() + if !running { + return } - }, - }) - - // 启动服务 - if err := r.rtmpServer.Serve(r.rtmpListener); err != nil { - if r.running { - log.Printf("❌ [RTMP] 服务异常: %v", err) + log.Printf("⚠️ [RTMP] Accept 失败: %v", err) + continue } + + go r.handleConnection(conn) } } -// RTMPHandler RTMP 连接处理器 -type RTMPHandler struct { - rtmp.DefaultHandler - server *RTMPServer - conn net.Conn - streamID string - isPublish bool +// handleConnection 处理 RTMP 连接 +func (r *RTMPServer) handleConnection(conn net.Conn) { + remoteAddr := conn.RemoteAddr().String() + log.Printf("🎬 [RTMP] 新连接: %s", remoteAddr) + + defer func() { + if err := recover(); err != nil { + log.Printf("❌ [RTMP] 连接处理 panic: %v", err) + } + conn.Close() + }() + + // 1. 执行握手 + if err := DoHandshake(conn, 30*time.Second); err != nil { + log.Printf("❌ [RTMP] 握手失败: %v", err) + return + } + + // 2. 创建读写器 + reader := NewChunkReader(conn) + writer := NewChunkWriter(conn) + + // 3. 创建连接处理器 + handler := &ConnectionHandler{ + server: r, + conn: conn, + reader: reader, + writer: writer, + remoteAddr: remoteAddr, + } + + // 4. 消息处理循环 + handler.serve() } -// OnServe 连接建立时调用 -func (h *RTMPHandler) OnServe(conn *rtmp.Conn) { - log.Printf("🎬 [RTMP] OnServe: %s", h.conn.RemoteAddr().String()) +// ConnectionHandler RTMP 连接处理器 +type ConnectionHandler struct { + server *RTMPServer + conn net.Conn + reader *ChunkReader + writer *ChunkWriter + remoteAddr string + + // 连接状态 + streamID string + isPublish bool + appName string + msgStreamID uint32 + + // 统计信息 + audioCount int64 + videoCount int64 + firstAudio bool + firstVideo bool } -// OnConnect 处理 connect 命令 -func (h *RTMPHandler) OnConnect(timestamp uint32, cmd *rtmpmsg.NetConnectionConnect) error { - log.Printf("🎬 [RTMP] OnConnect: app=%v", cmd.Command.App) +// serve 消息处理循环 +func (h *ConnectionHandler) serve() { + log.Printf("🎬 [RTMP] OnServe: %s", h.remoteAddr) + + msgCount := int64(0) + lastLogTime := time.Now() + + for { + msg, err := h.reader.ReadMessage() + if err != nil { + log.Printf("🔌 [RTMP] 连接关闭: %s (err: %v)", h.remoteAddr, err) + break + } + + msgCount++ + + // 每 100 条消息或每 5 秒记录一次统计信息 + if msgCount%100 == 0 || time.Since(lastLogTime) > 5*time.Second { + log.Printf("📊 [RTMP] 消息统计: stream=%s 总消息=%d audio=%d video=%d remote=%s", + h.streamID, msgCount, h.audioCount, h.videoCount, h.remoteAddr) + lastLogTime = time.Now() + } + + if err := h.handleMessage(msg); err != nil { + log.Printf("❌ [RTMP] 处理消息失败: %v", err) + break + } + } + + // 连接关闭时的清理 + h.onClose() +} + +// handleMessage 处理单个消息 +func (h *ConnectionHandler) handleMessage(msg *RTMPMessage) error { + switch msg.TypeID { + case RTMP_MSG_CHUNK_SIZE: + return h.handleSetChunkSize(msg) + case RTMP_MSG_ACK: + // 忽略 ACK + return nil + case RTMP_MSG_USER_CONTROL: + // 忽略 User Control + return nil + case RTMP_MSG_WIN_ACK_SIZE: + // 忽略 Window Ack Size + return nil + case RTMP_MSG_SET_PEER_BW: + // 忽略 Set Peer Bandwidth + return nil + case RTMP_MSG_AUDIO: + return h.handleAudio(msg) + case RTMP_MSG_VIDEO: + return h.handleVideo(msg) + case RTMP_MSG_AMF0_DATA: + return h.handleDataMessage(msg) + case RTMP_MSG_AMF0_CMD: + return h.handleCommand(msg) + default: + // 忽略未知消息 + return nil + } +} + +// handleSetChunkSize 处理 SetChunkSize 消息 +func (h *ConnectionHandler) handleSetChunkSize(msg *RTMPMessage) error { + if len(msg.Data) < 4 { + return fmt.Errorf("SetChunkSize 数据不足") + } + + newSize := binary.BigEndian.Uint32(msg.Data) + log.Printf("📝 [RTMP] SetChunkSize: %d -> %d (remote: %s)", h.reader.GetChunkSize(), newSize, h.remoteAddr) + + // 关键:立即更新 chunk size + h.reader.SetChunkSize(newSize) return nil } -// OnCreateStream 处理 createStream 命令 -func (h *RTMPHandler) OnCreateStream(timestamp uint32, cmd *rtmpmsg.NetConnectionCreateStream) error { - log.Printf("🎬 [RTMP] OnCreateStream") - return nil +// handleCommand 处理 AMF0 命令 +func (h *ConnectionHandler) handleCommand(msg *RTMPMessage) error { + values, err := DecodeAMF0(msg.Data) + if err != nil { + log.Printf("⚠️ [RTMP] 解析命令失败: %v", err) + return nil + } + + if len(values) == 0 { + return nil + } + + command, ok := values[0].(string) + if !ok { + return nil + } + + transactionID := float64(0) + if len(values) > 1 { + if tid, ok := values[1].(float64); ok { + transactionID = tid + } + } + + log.Printf("🎬 [RTMP] 命令: %s (tid: %.0f) from %s", command, transactionID, h.remoteAddr) + + switch command { + case "connect": + return h.handleConnect(values, transactionID) + case "releaseStream": + return h.handleReleaseStream(values) + case "FCPublish": + return h.handleFCPublish(values) + case "createStream": + return h.handleCreateStream(transactionID) + case "publish": + return h.handlePublish(values, msg.StreamID) + case "play": + return h.handlePlay(values, msg.StreamID) + case "deleteStream": + return h.handleDeleteStream(values) + case "FCUnpublish": + return h.handleFCUnpublish(values) + default: + // 忽略未知命令 + return nil + } } -// OnReleaseStream 处理 releaseStream 命令 -func (h *RTMPHandler) OnReleaseStream(timestamp uint32, cmd *rtmpmsg.NetConnectionReleaseStream) error { - streamName := cmd.StreamName +// handleConnect 处理 connect 命令 +func (h *ConnectionHandler) handleConnect(values []interface{}, tid float64) error { + // 解析 app 名称 + if len(values) > 2 { + if obj, ok := values[2].(AMF0Object); ok { + if app, ok := obj["app"].(string); ok { + h.appName = app + } + } + } + log.Printf("🎬 [RTMP] OnConnect: app=%s", h.appName) + + // 发送 Window Ack Size + if err := h.writer.WriteWindowAckSize(DEFAULT_WINDOW_SIZE); err != nil { + return err + } + + // 发送 Set Peer Bandwidth + if err := h.writer.WriteSetPeerBandwidth(DEFAULT_WINDOW_SIZE, 2); err != nil { + return err + } + + // 发送 Set Chunk Size + if err := h.writer.WriteSetChunkSize(4096); err != nil { + return err + } + + // 发送 _result + result := EncodeConnectResult(tid) + if err := h.writer.WriteCommand(3, 0, result); err != nil { + return err + } + + // 发送 onBWDone + bwDone := EncodeOnBWDone() + return h.writer.WriteCommand(3, 0, bwDone) +} + +// handleReleaseStream 处理 releaseStream 命令 +func (h *ConnectionHandler) handleReleaseStream(values []interface{}) error { + streamName := "" + if len(values) > 3 { + if name, ok := values[3].(string); ok { + streamName = name + } + } if idx := strings.Index(streamName, "?"); idx != -1 { streamName = streamName[:idx] } @@ -186,36 +373,59 @@ func (h *RTMPHandler) OnReleaseStream(timestamp uint32, cmd *rtmpmsg.NetConnecti return nil } -// OnDeleteStream 处理 deleteStream 命令 -func (h *RTMPHandler) OnDeleteStream(timestamp uint32, cmd *rtmpmsg.NetStreamDeleteStream) error { - log.Printf("🎬 [RTMP] OnDeleteStream") +// handleFCPublish 处理 FCPublish 命令 +func (h *ConnectionHandler) handleFCPublish(values []interface{}) error { + streamName := "" + if len(values) > 3 { + if name, ok := values[3].(string); ok { + streamName = name + } + } + log.Printf("🎬 [RTMP] OnFCPublish: %s", streamName) return nil } -// OnFCPublish 处理 FCPublish 命令 -func (h *RTMPHandler) OnFCPublish(timestamp uint32, cmd *rtmpmsg.NetStreamFCPublish) error { - log.Printf("🎬 [RTMP] OnFCPublish: %s", cmd.StreamName) - return nil +// handleCreateStream 处理 createStream 命令 +func (h *ConnectionHandler) handleCreateStream(tid float64) error { + log.Printf("🎬 [RTMP] OnCreateStream") + + h.msgStreamID = 1 + + // 发送 _result + result := EncodeCreateStreamResult(tid, float64(h.msgStreamID)) + return h.writer.WriteCommand(3, 0, result) } -// OnFCUnpublish 处理 FCUnpublish 命令 -func (h *RTMPHandler) OnFCUnpublish(timestamp uint32, cmd *rtmpmsg.NetStreamFCUnpublish) error { - log.Printf("🎬 [RTMP] OnFCUnpublish: %s", cmd.StreamName) - return nil -} +// handlePublish 处理 publish 命令 +func (h *ConnectionHandler) handlePublish(values []interface{}, streamID uint32) error { + publishingName := "" + publishingType := "" -// OnPublish 处理 publish 命令 -func (h *RTMPHandler) OnPublish(_ *rtmp.StreamContext, timestamp uint32, cmd *rtmpmsg.NetStreamPublish) error { - log.Printf("🎬 [RTMP] OnPublish: name=%s, type=%s", cmd.PublishingName, cmd.PublishingType) + if len(values) > 3 { + if name, ok := values[3].(string); ok { + publishingName = name + } + } + if len(values) > 4 { + if ptype, ok := values[4].(string); ok { + publishingType = ptype + } + } + + log.Printf("🎬 [RTMP] OnPublish: name=%s, type=%s, remote=%s", publishingName, publishingType, h.remoteAddr) // 解析 stream name,去掉 token 参数 - streamName := cmd.PublishingName + streamName := publishingName if idx := strings.Index(streamName, "?"); idx != -1 { streamName = streamName[:idx] } h.streamID = streamName h.isPublish = true + h.audioCount = 0 + h.videoCount = 0 + h.firstAudio = false + h.firstVideo = false log.Printf("🎬 [RTMP] 解析后的流ID: %s", h.streamID) @@ -226,7 +436,6 @@ func (h *RTMPHandler) OnPublish(_ *rtmp.StreamContext, timestamp uint32, cmd *rt if !exists { log.Printf("⚠️ [RTMP] 流未注册,自动创建: %s", h.streamID) - // 自动创建流(用于测试) stream = &RTMPStream{ ID: h.streamID, CreatedAt: time.Now(), @@ -243,18 +452,40 @@ func (h *RTMPHandler) OnPublish(_ *rtmp.StreamContext, timestamp uint32, cmd *rt stream.mu.Lock() stream.IsActive = true + // 重置流数据 + stream.audioHeader = nil + stream.videoHeader = nil + stream.metaData = nil + stream.gopCache = make([][]byte, 0) stream.mu.Unlock() - log.Printf("✅ [RTMP] 开始推流: %s", h.streamID) + // 发送 Stream Begin + if err := h.writer.WriteStreamBegin(streamID); err != nil { + return err + } + // 发送 onStatus (NetStream.Publish.Start) + status := EncodeOnStatus("NetStream.Publish.Start", "status", "Publishing started.") + if err := h.writer.WriteCommand(5, streamID, status); err != nil { + return err + } + + log.Printf("✅ [RTMP] 开始推流: %s (等待音视频数据...)", h.streamID) return nil } -// OnPlay 处理 play 命令 -func (h *RTMPHandler) OnPlay(ctx *rtmp.StreamContext, timestamp uint32, cmd *rtmpmsg.NetStreamPlay) error { - log.Printf("🎬 [RTMP] OnPlay: name=%s", cmd.StreamName) +// handlePlay 处理 play 命令 +func (h *ConnectionHandler) handlePlay(values []interface{}, streamID uint32) error { + streamName := "" + if len(values) > 3 { + if name, ok := values[3].(string); ok { + streamName = name + } + } - h.streamID = cmd.StreamName + log.Printf("🎬 [RTMP] OnPlay: name=%s", streamName) + + h.streamID = streamName h.isPublish = false // 检查流是否存在 @@ -264,19 +495,50 @@ func (h *RTMPHandler) OnPlay(ctx *rtmp.StreamContext, timestamp uint32, cmd *rtm if !exists || !stream.IsActive { log.Printf("⚠️ [RTMP] 流不存在或未激活: %s", h.streamID) - return fmt.Errorf("stream not found: %s", h.streamID) + status := EncodeOnStatus("NetStream.Play.StreamNotFound", "error", "Stream not found.") + return h.writer.WriteCommand(5, streamID, status) + } + + // 发送 Stream Begin + if err := h.writer.WriteStreamBegin(streamID); err != nil { + return err + } + + // 发送 onStatus (NetStream.Play.Start) + status := EncodeOnStatus("NetStream.Play.Start", "status", "Playing started.") + if err := h.writer.WriteCommand(5, streamID, status); err != nil { + return err } log.Printf("✅ [RTMP] 开始播放: %s", h.streamID) - return nil } -// OnSetDataFrame 处理 metadata -func (h *RTMPHandler) OnSetDataFrame(timestamp uint32, data *rtmpmsg.NetStreamSetDataFrame) error { - log.Printf("🎬 [RTMP] OnSetDataFrame") +// handleDeleteStream 处理 deleteStream 命令 +func (h *ConnectionHandler) handleDeleteStream(values []interface{}) error { + log.Printf("🎬 [RTMP] OnDeleteStream") + return nil +} + +// handleFCUnpublish 处理 FCUnpublish 命令 +func (h *ConnectionHandler) handleFCUnpublish(values []interface{}) error { + streamName := "" + if len(values) > 3 { + if name, ok := values[3].(string); ok { + streamName = name + } + } + log.Printf("🎬 [RTMP] OnFCUnpublish: %s", streamName) + return nil +} + +// handleDataMessage 处理数据消息 (@setDataFrame) +func (h *ConnectionHandler) handleDataMessage(msg *RTMPMessage) error { + log.Printf("🎬 [RTMP] OnSetDataFrame: stream=%s payloadLen=%d audioCount=%d videoCount=%d", + h.streamID, len(msg.Data), h.audioCount, h.videoCount) if h.streamID == "" { + log.Printf("⚠️ [RTMP] OnSetDataFrame: streamID 为空") return nil } @@ -284,31 +546,39 @@ func (h *RTMPHandler) OnSetDataFrame(timestamp uint32, data *rtmpmsg.NetStreamSe stream, exists := h.server.streams[h.streamID] h.server.streamsMu.RUnlock() - if exists && len(data.Payload) > 0 { + if exists && len(msg.Data) > 0 { stream.mu.Lock() - stream.metaData = data.Payload + stream.metaData = make([]byte, len(msg.Data)) + copy(stream.metaData, msg.Data) stream.mu.Unlock() + log.Printf("✅ [RTMP] OnSetDataFrame: 已保存 metadata") } return nil } -// OnAudio 处理音频数据 -func (h *RTMPHandler) OnAudio(timestamp uint32, payload io.Reader) error { +// handleAudio 处理音频数据 +func (h *ConnectionHandler) handleAudio(msg *RTMPMessage) error { + h.audioCount++ + + if !h.firstAudio { + h.firstAudio = true + log.Printf("🎵 [RTMP] OnAudio 首帧: stream=%s timestamp=%d isPublish=%v", h.streamID, msg.Timestamp, h.isPublish) + } + if h.streamID == "" || !h.isPublish { return nil } - // 读取音频数据 - data, err := io.ReadAll(payload) - if err != nil { - return err - } - + data := msg.Data if len(data) == 0 { return nil } + if h.audioCount%100 == 0 { + log.Printf("🎵 [RTMP] OnAudio 统计: stream=%s audioCount=%d videoCount=%d", h.streamID, h.audioCount, h.videoCount) + } + h.server.streamsMu.RLock() stream, exists := h.server.streams[h.streamID] h.server.streamsMu.RUnlock() @@ -332,30 +602,38 @@ func (h *RTMPHandler) OnAudio(timestamp uint32, payload io.Reader) error { } // 创建 FLV audio tag - flvData := h.createFLVAudioTag(timestamp, data) - if flvData != nil { - h.broadcastToSubscribers(h.streamID, flvData) - } + flvData := createFLVTag(8, msg.Timestamp, data) + h.broadcastToSubscribers(h.streamID, flvData) return nil } -// OnVideo 处理视频数据 -func (h *RTMPHandler) OnVideo(timestamp uint32, payload io.Reader) error { +// handleVideo 处理视频数据 +func (h *ConnectionHandler) handleVideo(msg *RTMPMessage) error { + h.videoCount++ + + if !h.firstVideo { + h.firstVideo = true + log.Printf("🎥 [RTMP] OnVideo 首帧: stream=%s timestamp=%d isPublish=%v", h.streamID, msg.Timestamp, h.isPublish) + } + if h.streamID == "" || !h.isPublish { return nil } - // 读取视频数据 - data, err := io.ReadAll(payload) - if err != nil { - return err - } - + data := msg.Data if len(data) == 0 { return nil } + if h.videoCount == 1 { + log.Printf("🎥 [RTMP] OnVideo 首帧数据: stream=%s len=%d", h.streamID, len(data)) + } + + if h.videoCount%30 == 0 { + log.Printf("🎥 [RTMP] OnVideo 统计: stream=%s audioCount=%d videoCount=%d", h.streamID, h.audioCount, h.videoCount) + } + h.server.streamsMu.RLock() stream, exists := h.server.streams[h.streamID] h.server.streamsMu.RUnlock() @@ -387,131 +665,23 @@ func (h *RTMPHandler) OnVideo(timestamp uint32, payload io.Reader) error { } // 创建 FLV video tag - flvData := h.createFLVVideoTag(timestamp, data) - if flvData != nil { - // 缓存到 GOP - stream.mu.Lock() - stream.gopCache = append(stream.gopCache, flvData) - // 限制 GOP 缓存大小 - if len(stream.gopCache) > 300 { - stream.gopCache = stream.gopCache[len(stream.gopCache)-300:] - } - stream.mu.Unlock() + flvData := createFLVTag(9, msg.Timestamp, data) - h.broadcastToSubscribers(h.streamID, flvData) + // 缓存到 GOP + stream.mu.Lock() + stream.gopCache = append(stream.gopCache, flvData) + if len(stream.gopCache) > 300 { + stream.gopCache = stream.gopCache[len(stream.gopCache)-300:] } + stream.mu.Unlock() + + h.broadcastToSubscribers(h.streamID, flvData) return nil } -// OnUnknownMessage 处理未知消息 -func (h *RTMPHandler) OnUnknownMessage(timestamp uint32, msg rtmpmsg.Message) error { - return nil -} - -// OnUnknownCommandMessage 处理未知命令消息 -func (h *RTMPHandler) OnUnknownCommandMessage(timestamp uint32, cmd *rtmpmsg.CommandMessage) error { - return nil -} - -// OnUnknownDataMessage 处理未知数据消息 -func (h *RTMPHandler) OnUnknownDataMessage(timestamp uint32, data *rtmpmsg.DataMessage) error { - return nil -} - -// createFLVAudioTag 创建 FLV 音频 tag -func (h *RTMPHandler) createFLVAudioTag(timestamp uint32, data []byte) []byte { - // FLV tag 格式: - // TagType (1 byte): 8 = audio - // DataSize (3 bytes): big-endian - // Timestamp (3 bytes): big-endian - // TimestampExtended (1 byte) - // StreamID (3 bytes): always 0 - // Data - // PreviousTagSize (4 bytes): big-endian - - dataSize := len(data) - tagSize := 11 + dataSize + 4 // header + data + previous tag size - - tag := make([]byte, tagSize) - - // Tag type - tag[0] = 8 // Audio - - // 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 (11 + dataSize) - 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 -} - -// createFLVVideoTag 创建 FLV 视频 tag -func (h *RTMPHandler) createFLVVideoTag(timestamp uint32, data []byte) []byte { - dataSize := len(data) - tagSize := 11 + dataSize + 4 - - tag := make([]byte, tagSize) - - // Tag type - tag[0] = 9 // Video - - // 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 -} - // broadcastToSubscribers 向所有订阅者广播数据 -func (h *RTMPHandler) broadcastToSubscribers(streamID string, data []byte) { +func (h *ConnectionHandler) broadcastToSubscribers(streamID string, data []byte) { h.server.subMu.RLock() subs, exists := h.server.subscribers[streamID] if !exists { @@ -529,9 +699,15 @@ func (h *RTMPHandler) broadcastToSubscribers(streamID string, data []byte) { h.server.subMu.RUnlock() } -// OnClose 连接关闭 -func (h *RTMPHandler) OnClose() { - log.Printf("🎬 [RTMP] OnClose: stream=%s, publish=%v", h.streamID, h.isPublish) +// onClose 连接关闭 +func (h *ConnectionHandler) onClose() { + log.Printf("🔌 [RTMP] OnClose: stream=%s, publish=%v, remote=%s, audioCount=%d, videoCount=%d", + h.streamID, h.isPublish, h.remoteAddr, h.audioCount, h.videoCount) + + if h.isPublish && h.audioCount > 0 && h.videoCount == 0 { + log.Printf("⚠️ [RTMP] 警告: 只收到音频(%d帧),没有收到视频!stream=%s", h.audioCount, h.streamID) + log.Printf("⚠️ [RTMP] 可能原因: 1.推流端未启用视频 2.视频编码格式不支持 3.连接过早断开") + } if h.isPublish && h.streamID != "" { h.server.streamsMu.Lock() @@ -541,145 +717,12 @@ func (h *RTMPHandler) OnClose() { stream.mu.Unlock() } h.server.streamsMu.Unlock() - log.Printf("⏹️ [RTMP] 停止推流: %s", h.streamID) + log.Printf("⏹️ [RTMP] 停止推流: %s (总计: 音频%d帧, 视频%d帧)", h.streamID, h.audioCount, h.videoCount) } } -// startHTTPFLVServer 启动 HTTP-FLV 服务 -func (r *RTMPServer) startHTTPFLVServer() { - addr := fmt.Sprintf(":%d", r.config.HTTPFLVPort) - - mux := http.NewServeMux() - mux.HandleFunc("/live/", r.handleFLVRequest) - mux.HandleFunc("/api/streams", r.handleStreamsAPI) - - r.httpServer = &http.Server{ - Addr: addr, - Handler: mux, - } - - log.Printf("🎬 [RTMP] HTTP-FLV 服务监听: %s", addr) - - if err := r.httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed { - log.Printf("❌ [RTMP] HTTP-FLV 启动失败: %v", err) - } -} - -// handleFLVRequest 处理 FLV 请求 -func (r *RTMPServer) handleFLVRequest(w http.ResponseWriter, req *http.Request) { - // 解析流ID: /live/{streamID}.flv - path := req.URL.Path - if len(path) < 7 { - http.Error(w, "Invalid path", http.StatusBadRequest) - return - } - - // 提取 streamID - streamPath := path[6:] // 去掉 "/live/" - if len(streamPath) > 4 && streamPath[len(streamPath)-4:] == ".flv" { - streamPath = streamPath[:len(streamPath)-4] - } - - r.streamsMu.RLock() - stream, exists := r.streams[streamPath] - r.streamsMu.RUnlock() - - if !exists || !stream.IsActive { - http.Error(w, "Stream not found", http.StatusNotFound) - return - } - - log.Printf("🎬 [RTMP] FLV 请求: %s", streamPath) - - // 设置响应头 - w.Header().Set("Content-Type", "video/x-flv") - w.Header().Set("Access-Control-Allow-Origin", "*") - w.Header().Set("Transfer-Encoding", "chunked") - w.Header().Set("Connection", "keep-alive") - w.Header().Set("Cache-Control", "no-cache") - - // 发送 FLV header - // FLV signature (3 bytes) + version (1 byte) + flags (1 byte) + header size (4 bytes) - flvHeader := []byte{0x46, 0x4C, 0x56, 0x01, 0x05, 0x00, 0x00, 0x00, 0x09, 0x00, 0x00, 0x00, 0x00} - w.Write(flvHeader) - - if f, ok := w.(http.Flusher); ok { - f.Flush() - } - - // 创建订阅者 - subID := fmt.Sprintf("%d", time.Now().UnixNano()) - sub := &Subscriber{ - ID: subID, - StreamID: streamPath, - DataChan: make(chan []byte, 100), - Done: make(chan struct{}), - } - - // 注册订阅者 - r.subMu.Lock() - if r.subscribers[streamPath] == nil { - r.subscribers[streamPath] = make(map[string]*Subscriber) - } - r.subscribers[streamPath][subID] = sub - r.subMu.Unlock() - - // 清理 - defer func() { - r.subMu.Lock() - delete(r.subscribers[streamPath], subID) - r.subMu.Unlock() - close(sub.Done) - }() - - // 发送缓存的头信息 - stream.mu.RLock() - - // 发送 metadata - if len(stream.metaData) > 0 { - w.Write(stream.metaData) - } - - // 发送视频 sequence header - if len(stream.videoHeader) > 0 { - flvTag := createFLVTagFromData(9, 0, stream.videoHeader) - w.Write(flvTag) - } - - // 发送音频 sequence header - if len(stream.audioHeader) > 0 { - flvTag := createFLVTagFromData(8, 0, stream.audioHeader) - w.Write(flvTag) - } - - // 发送 GOP 缓存 - for _, data := range stream.gopCache { - w.Write(data) - } - stream.mu.RUnlock() - - if f, ok := w.(http.Flusher); ok { - f.Flush() - } - - // 持续发送数据 - for { - select { - case <-req.Context().Done(): - return - case data := <-sub.DataChan: - if _, err := w.Write(data); err != nil { - return - } - if f, ok := w.(http.Flusher); ok { - f.Flush() - } - } - } -} - -// createFLVTagFromData 从原始数据创建 FLV tag -func createFLVTagFromData(tagType byte, timestamp uint32, data []byte) []byte { +// createFLVTag 创建 FLV tag +func createFLVTag(tagType byte, timestamp uint32, data []byte) []byte { dataSize := len(data) tagSize := 11 + dataSize + 4 @@ -708,6 +751,174 @@ func createFLVTagFromData(tagType byte, timestamp uint32, data []byte) []byte { return tag } +// startHTTPFLVServer 启动 HTTP-FLV 服务 +func (r *RTMPServer) startHTTPFLVServer() { + addr := fmt.Sprintf(":%d", r.config.HTTPFLVPort) + + mux := http.NewServeMux() + mux.HandleFunc("/live/", r.handleFLVRequest) + mux.HandleFunc("/api/streams", r.handleStreamsAPI) + + r.httpServer = &http.Server{ + Addr: addr, + Handler: mux, + } + + log.Printf("🎬 [RTMP] HTTP-FLV 服务监听: %s", addr) + + if err := r.httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed { + log.Printf("❌ [RTMP] HTTP-FLV 启动失败: %v", err) + } +} + +// handleFLVRequest 处理 FLV 请求 +func (r *RTMPServer) handleFLVRequest(w http.ResponseWriter, req *http.Request) { + path := req.URL.Path + if len(path) < 7 { + http.Error(w, "Invalid path", http.StatusBadRequest) + return + } + + streamPath := path[6:] // 去掉 "/live/" + if len(streamPath) > 4 && streamPath[len(streamPath)-4:] == ".flv" { + streamPath = streamPath[:len(streamPath)-4] + } + + log.Printf("🎬 [RTMP] FLV 请求: %s (from: %s)", streamPath, req.RemoteAddr) + + // 等待流就绪 + var stream *RTMPStream + var exists bool + maxWait := 10 * time.Second + waitInterval := 500 * time.Millisecond + waited := time.Duration(0) + + for waited < maxWait { + r.streamsMu.RLock() + stream, exists = r.streams[streamPath] + r.streamsMu.RUnlock() + + if exists && stream.IsActive { + break + } + + select { + case <-req.Context().Done(): + log.Printf("⚠️ [RTMP] FLV 请求已取消: %s", streamPath) + return + default: + } + + if waited == 0 { + log.Printf("⏳ [RTMP] 等待流就绪: %s", streamPath) + } + + time.Sleep(waitInterval) + waited += waitInterval + } + + if !exists || !stream.IsActive { + log.Printf("❌ [RTMP] 流不存在或未激活: %s (waited: %v)", streamPath, waited) + http.Error(w, "Stream not found or not active", http.StatusNotFound) + return + } + + log.Printf("✅ [RTMP] 流已就绪,开始 FLV 传输: %s (waited: %v)", streamPath, waited) + + // 设置响应头 + w.Header().Set("Content-Type", "video/x-flv") + w.Header().Set("Access-Control-Allow-Origin", "*") + w.Header().Set("Transfer-Encoding", "chunked") + w.Header().Set("Connection", "keep-alive") + w.Header().Set("Cache-Control", "no-cache") + + // 发送 FLV header + flvHeader := []byte{0x46, 0x4C, 0x56, 0x01, 0x05, 0x00, 0x00, 0x00, 0x09, 0x00, 0x00, 0x00, 0x00} + w.Write(flvHeader) + + if f, ok := w.(http.Flusher); ok { + f.Flush() + } + + // 创建订阅者 + subID := fmt.Sprintf("%d", time.Now().UnixNano()) + sub := &Subscriber{ + ID: subID, + StreamID: streamPath, + DataChan: make(chan []byte, 100), + Done: make(chan struct{}), + } + + // 注册订阅者 + r.subMu.Lock() + if r.subscribers[streamPath] == nil { + r.subscribers[streamPath] = make(map[string]*Subscriber) + } + r.subscribers[streamPath][subID] = sub + r.subMu.Unlock() + + defer func() { + r.subMu.Lock() + delete(r.subscribers[streamPath], subID) + r.subMu.Unlock() + close(sub.Done) + }() + + // 发送缓存的头信息 + stream.mu.RLock() + hasMetadata := len(stream.metaData) > 0 + hasVideoHeader := len(stream.videoHeader) > 0 + hasAudioHeader := len(stream.audioHeader) > 0 + gopLen := len(stream.gopCache) + + log.Printf("📦 [RTMP] FLV 缓存状态: stream=%s metadata=%v videoHeader=%v audioHeader=%v gopLen=%d", + streamPath, hasMetadata, hasVideoHeader, hasAudioHeader, gopLen) + + if hasMetadata { + w.Write(stream.metaData) + } + + if hasVideoHeader { + flvTag := createFLVTag(9, 0, stream.videoHeader) + w.Write(flvTag) + } + + if hasAudioHeader { + flvTag := createFLVTag(8, 0, stream.audioHeader) + w.Write(flvTag) + } + + for _, data := range stream.gopCache { + w.Write(data) + } + stream.mu.RUnlock() + + if f, ok := w.(http.Flusher); ok { + f.Flush() + } + + log.Printf("▶️ [RTMP] FLV 开始实时传输: stream=%s subscriber=%s", streamPath, subID) + + // 持续发送数据 + dataCount := 0 + for { + select { + case <-req.Context().Done(): + log.Printf("⏹️ [RTMP] FLV 传输结束 (客户端断开): stream=%s dataCount=%d", streamPath, dataCount) + return + case data := <-sub.DataChan: + if _, err := w.Write(data); err != nil { + log.Printf("⏹️ [RTMP] FLV 传输结束 (写入失败): stream=%s dataCount=%d err=%v", streamPath, dataCount, err) + return + } + dataCount++ + if f, ok := w.(http.Flusher); ok { + f.Flush() + } + } + } +} + // handleStreamsAPI 处理流列表 API func (r *RTMPServer) handleStreamsAPI(w http.ResponseWriter, req *http.Request) { w.Header().Set("Content-Type", "application/json") @@ -739,11 +950,6 @@ func (r *RTMPServer) Stop() { r.running = false - // 关闭 RTMP 服务器 - if r.rtmpServer != nil { - r.rtmpServer.Close() - } - // 关闭 RTMP 监听器 if r.rtmpListener != nil { r.rtmpListener.Close() @@ -791,19 +997,15 @@ func (r *RTMPServer) GenerateStreamURLs(roomID, userID string) (*RTMPStream, err } r.mu.RUnlock() - // 生成唯一的流ID streamID := fmt.Sprintf("%s_%s_%d", roomID, userID, time.Now().UnixNano()) - // 获取服务器配置 publicIP := r.config.PublicIP if publicIP == "" { publicIP = viper.GetString("turn.public_ip") } - // 生成鉴权 Token token := generateStreamToken(streamID) - // 生成 URL stream := &RTMPStream{ ID: streamID, RoomID: roomID, @@ -879,7 +1081,6 @@ func (r *RTMPServer) RemoveStream(streamID string) { log.Printf("🎬 [RTMP] 移除流 | Stream:%s", streamID) } - // 移除订阅者 r.subMu.Lock() delete(r.subscribers, streamID) r.subMu.Unlock() @@ -903,7 +1104,6 @@ func (r *RTMPServer) RemoveStreamsByUser(userID string) { delete(r.streams, streamID) log.Printf("🎬 [RTMP] 移除用户流 | User:%s Stream:%s", userID, streamID) - // 移除订阅者 r.subMu.Lock() delete(r.subscribers, streamID) r.subMu.Unlock() @@ -929,7 +1129,6 @@ func (r *RTMPServer) RemoveStreamsByRoom(roomID string) { delete(r.streams, streamID) log.Printf("🎬 [RTMP] 移除房间流 | Room:%s Stream:%s", roomID, streamID) - // 移除订阅者 r.subMu.Lock() delete(r.subscribers, streamID) r.subMu.Unlock() @@ -969,7 +1168,7 @@ func (r *RTMPServer) cleanupExpiredStreams() { r.streamsMu.Lock() defer r.streamsMu.Unlock() - expireTime := time.Now().Add(-1 * time.Hour) // 1小时过期 + expireTime := time.Now().Add(-1 * time.Hour) for streamID, stream := range r.streams { if stream.CreatedAt.Before(expireTime) && !stream.IsActive { @@ -997,15 +1196,13 @@ func generateStreamToken(streamID string) string { return base64.URLEncoding.EncodeToString(mac.Sum(nil)) } +// GenerateStreamToken 生成流鉴权 Token (公开接口) +func GenerateStreamToken(streamID string) string { + return generateStreamToken(streamID) +} + // ValidateStreamToken 验证流鉴权 Token func ValidateStreamToken(streamID, token string) bool { expectedToken := generateStreamToken(streamID) return token == expectedToken } - -// 确保不使用的导入被使用 -var ( - _ = bytes.NewReader - _ = flv.NewEncoder - _ = flvtag.TagTypeVideo -) diff --git a/internal/mediaserver/rtmp_amf0.go b/internal/mediaserver/rtmp_amf0.go new file mode 100644 index 0000000..729fdcd --- /dev/null +++ b/internal/mediaserver/rtmp_amf0.go @@ -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() +} + + diff --git a/internal/mediaserver/rtmp_handshake.go b/internal/mediaserver/rtmp_handshake.go new file mode 100644 index 0000000..424cb15 --- /dev/null +++ b/internal/mediaserver/rtmp_handshake.go @@ -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 +} + + diff --git a/internal/mediaserver/rtmp_protocol.go b/internal/mediaserver/rtmp_protocol.go new file mode 100644 index 0000000..d8f628e --- /dev/null +++ b/internal/mediaserver/rtmp_protocol.go @@ -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 必须有有效的 prevHeader(MessageLength 从 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, + }) +} + diff --git a/internal/mediaserver/server.go b/internal/mediaserver/server.go index 49e520f..887c0c7 100644 --- a/internal/mediaserver/server.go +++ b/internal/mediaserver/server.go @@ -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() diff --git a/internal/mediaserver/ws_rtmp_proxy.go b/internal/mediaserver/ws_rtmp_proxy.go new file mode 100644 index 0000000..3a2fc6d --- /dev/null +++ b/internal/mediaserver/ws_rtmp_proxy.go @@ -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) + diff --git a/internal/middleware/auth.go b/internal/middleware/auth.go index 67e3720..b77fced 100644 --- a/internal/middleware/auth.go +++ b/internal/middleware/auth.go @@ -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) diff --git a/internal/middleware/request_log.go b/internal/middleware/request_log.go index 7052c35..b88191b 100644 --- a/internal/middleware/request_log.go +++ b/internal/middleware/request_log.go @@ -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)