Skip to content
This repository was archived by the owner on Jul 10, 2026. It is now read-only.

Commit 136f764

Browse files
committed
feat(ws): improve mirroring handler with setup validation
1 parent 1cb38ab commit 136f764

1 file changed

Lines changed: 107 additions & 42 deletions

File tree

internal/handler/ws/mirroring_handler.go

Lines changed: 107 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package ws
22

33
import (
44
"context"
5+
"encoding/json"
56
"net/http"
67
"sync"
78
"time"
@@ -34,6 +35,7 @@ func NewMirroringHandler(mirroringService mirroring_service.IMirroringService, l
3435
},
3536
Logger: logger,
3637
activeConns: make(map[string]*websocket.Conn),
38+
Validator: validator,
3739
}
3840
}
3941

@@ -92,26 +94,35 @@ func (h *MirroringHandler) StartMirroringStream(w http.ResponseWriter, r *http.R
9294
zap.String("serial", serial),
9395
)
9496

95-
if !h.MirroringService.IsRunning(serial) {
96-
var setup model.MirroringSetupRequest
97-
client, err = h.MirroringService.StartMirroring(ctx, serial, &setup)
98-
if err != nil {
99-
h.Logger.Error("failed to start mirroring",
100-
zap.String("serial", serial),
101-
zap.Error(err),
102-
)
103-
errorMsg := map[string]interface{}{
104-
"type": "error",
105-
"message": "Failed to start mirroring: " + err.Error(),
106-
}
107-
if err := conn.WriteJSON(errorMsg); err != nil {
108-
h.Logger.Error("failed to send error message",
97+
setupReceived := make(chan *model.MirroringSetupRequest, 1)
98+
99+
go h.handleWebSocketMessages(ctx, cancel, conn, serial, setupReceived)
100+
101+
var setup *model.MirroringSetupRequest
102+
select {
103+
case setup = <-setupReceived:
104+
if !h.MirroringService.IsRunning(serial) {
105+
client, err = h.MirroringService.StartMirroring(ctx, serial, setup)
106+
if err != nil {
107+
h.Logger.Error("failed to start mirroring",
109108
zap.String("serial", serial),
110109
zap.Error(err),
111110
)
111+
errorMsg := map[string]interface{}{
112+
"type": "error",
113+
"message": "Failed to start mirroring: " + err.Error(),
114+
}
115+
if err := conn.WriteJSON(errorMsg); err != nil {
116+
h.Logger.Error("failed to send error message",
117+
zap.String("serial", serial),
118+
zap.Error(err),
119+
)
120+
}
121+
return
112122
}
113-
return
114123
}
124+
case <-ctx.Done():
125+
return
115126
}
116127

117128
successMsg := map[string]interface{}{
@@ -132,7 +143,6 @@ func (h *MirroringHandler) StartMirroringStream(w http.ResponseWriter, r *http.R
132143
pingTicker := time.NewTicker(30 * time.Second)
133144
defer pingTicker.Stop()
134145

135-
go h.handleWebSocketControl(ctx, cancel, conn, serial)
136146
handleVideoChunk := func(chunk []byte) {
137147
if len(chunk) == 0 {
138148
return
@@ -187,38 +197,93 @@ func (h *MirroringHandler) StartMirroringStream(w http.ResponseWriter, r *http.R
187197
}
188198
}
189199

190-
func (h *MirroringHandler) handleWebSocketControl(ctx context.Context, cancelCtx context.CancelFunc, conn *websocket.Conn, serial string) {
200+
func (h *MirroringHandler) handleWebSocketMessages(ctx context.Context, cancelCtx context.CancelFunc, conn *websocket.Conn, serial string, setupChan chan<- *model.MirroringSetupRequest) {
191201
conn.SetReadLimit(1024)
192202

193-
go func() {
194-
for {
195-
select {
196-
case <-ctx.Done():
203+
for {
204+
select {
205+
case <-ctx.Done():
206+
return
207+
default:
208+
_, messageBytes, err := conn.ReadMessage()
209+
if err != nil {
210+
h.Logger.Info("client disconnected",
211+
zap.String("serial", serial),
212+
zap.Error(err),
213+
)
214+
cancelCtx()
197215
return
198-
default:
199-
_, messageBytes, err := conn.ReadMessage()
200-
if err != nil {
201-
h.Logger.Info("client disconnected",
202-
zap.String("serial", serial),
203-
zap.Error(err),
204-
)
205-
cancelCtx()
206-
return
207-
}
216+
}
208217

209-
if len(messageBytes) > 1024 {
210-
h.Logger.Warn("received oversized message, ignoring",
211-
zap.String("serial", serial),
212-
zap.Int("size", len(messageBytes)),
213-
)
214-
continue
215-
}
218+
if len(messageBytes) > 1024 {
219+
h.Logger.Warn("received oversized message, ignoring",
220+
zap.String("serial", serial),
221+
zap.Int("size", len(messageBytes)),
222+
)
223+
continue
224+
}
225+
226+
if len(messageBytes) > 0 {
227+
var message map[string]interface{}
228+
if err := json.Unmarshal(messageBytes, &message); err == nil {
229+
if msgType, ok := message["type"].(string); ok && msgType == "setup" {
230+
var setupRequest model.MirroringSetupRequest
231+
if err := json.Unmarshal(messageBytes, &setupRequest); err != nil {
232+
h.Logger.Error("failed to parse setup message",
233+
zap.String("serial", serial),
234+
zap.Error(err),
235+
)
236+
errorMsg := map[string]interface{}{
237+
"type": "error",
238+
"message": "Invalid setup message format",
239+
}
240+
if err := conn.WriteJSON(errorMsg); err != nil {
241+
h.Logger.Error("failed to send error message",
242+
zap.String("serial", serial),
243+
zap.Error(err),
244+
)
245+
}
246+
continue
247+
}
248+
249+
if err := h.Validator.Struct(&setupRequest); err != nil {
250+
h.Logger.Error("setup validation failed",
251+
zap.String("serial", serial),
252+
zap.Error(err),
253+
)
254+
errorMsg := map[string]interface{}{
255+
"type": "error",
256+
"message": "Setup validation failed: " + err.Error(),
257+
}
258+
if err := conn.WriteJSON(errorMsg); err != nil {
259+
h.Logger.Error("failed to send validation error message",
260+
zap.String("serial", serial),
261+
zap.Error(err),
262+
)
263+
}
264+
continue
265+
}
216266

217-
if len(messageBytes) > 0 {
218-
h.MirroringService.HandleControlMessage(serial, messageBytes)
267+
h.Logger.Info("received valid setup configuration",
268+
zap.String("serial", serial),
269+
zap.Uint8("fps", setupRequest.FPS),
270+
zap.Uint32("bitrate", setupRequest.Bitrate),
271+
zap.Uint16("resolution", setupRequest.Resolution),
272+
)
273+
274+
select {
275+
case setupChan <- &setupRequest:
276+
default:
277+
h.Logger.Warn("setup channel full, ignoring duplicate setup message",
278+
zap.String("serial", serial),
279+
)
280+
}
281+
continue
282+
}
219283
}
284+
285+
h.MirroringService.HandleControlMessage(serial, messageBytes)
220286
}
221287
}
222-
}()
223-
288+
}
224289
}

0 commit comments

Comments
 (0)