diff --git a/pkg/service/agentservice.go b/pkg/service/agentservice.go index 8a7a6160f..f2ba672b2 100644 --- a/pkg/service/agentservice.go +++ b/pkg/service/agentservice.go @@ -198,10 +198,27 @@ func (s *AgentHandler) HandleConnection(conn *websocket.Conn) { } func (s *AgentHandler) handleRegister(worker *worker, msg *livekit.RegisterWorkerRequest) { + if err := s.doHandleRegister(worker, msg); err != nil { + logger.Errorw("failed to register worker", err, "workerID", msg.WorkerId, "jobType", msg.Type) + worker.conn.Close() + } +} + +func (s *AgentHandler) doHandleRegister(worker *worker, msg *livekit.RegisterWorkerRequest) error { + if msg.WorkerId == "" { + return errors.New("invalid worker id") + } + s.mu.Lock() + if worker.id != "" { + s.mu.Unlock() + return errors.New("worker already registered") + } + switch msg.Type { case livekit.JobType_JT_ROOM: worker.id = msg.WorkerId + worker.jobType = msg.Type delete(s.unregistered, worker.conn) s.roomWorkers[worker.id] = worker @@ -216,6 +233,7 @@ func (s *AgentHandler) handleRegister(worker *worker, msg *livekit.RegisterWorke case livekit.JobType_JT_PUBLISHER: worker.id = msg.WorkerId + worker.jobType = msg.Type delete(s.unregistered, worker.conn) s.publisherWorkers[worker.id] = worker @@ -227,6 +245,9 @@ func (s *AgentHandler) handleRegister(worker *worker, msg *livekit.RegisterWorke s.publisherRegistered = true } } + default: + s.mu.Unlock() + return errors.New("invalid job type") } s.mu.Unlock() @@ -241,6 +262,8 @@ func (s *AgentHandler) handleRegister(worker *worker, msg *livekit.RegisterWorke if err != nil { logger.Errorw("failed to write server message", err) } + + return nil } func (s *AgentHandler) handleAvailability(w *worker, msg *livekit.AvailabilityResponse) {