diff --git a/go.mod b/go.mod index e0d1acf60..81b773866 100644 --- a/go.mod +++ b/go.mod @@ -19,7 +19,7 @@ require ( github.com/jxskiss/base62 v1.1.0 github.com/livekit/mageutil v0.0.0-20230125210925-54e8a70427c1 github.com/livekit/mediatransportutil v0.0.0-20240730083616-559fa5ece598 - github.com/livekit/protocol v1.21.1-0.20240913074525-1f5de7d620c4 + github.com/livekit/protocol v1.21.1-0.20240918070421-da4a63076085 github.com/livekit/psrpc v0.5.3-0.20240616012458-ac39c8549a0a github.com/mackerelio/go-osstat v0.2.5 github.com/magefile/mage v1.15.0 diff --git a/go.sum b/go.sum index 4002f6667..4052e956d 100644 --- a/go.sum +++ b/go.sum @@ -169,8 +169,8 @@ github.com/livekit/mageutil v0.0.0-20230125210925-54e8a70427c1 h1:jm09419p0lqTkD github.com/livekit/mageutil v0.0.0-20230125210925-54e8a70427c1/go.mod h1:Rs3MhFwutWhGwmY1VQsygw28z5bWcnEYmS1OG9OxjOQ= github.com/livekit/mediatransportutil v0.0.0-20240730083616-559fa5ece598 h1:yLlkHk2feSLHstD9n4VKg7YEBR4rLODTI4WE8gNBEnQ= github.com/livekit/mediatransportutil v0.0.0-20240730083616-559fa5ece598/go.mod h1:jwKUCmObuiEDH0iiuJHaGMXwRs3RjrB4G6qqgkr/5oE= -github.com/livekit/protocol v1.21.1-0.20240913074525-1f5de7d620c4 h1:JhLovaAv+UxOTYpjQxN3IRV3kuJFb9f3LZmg/KqvHkc= -github.com/livekit/protocol v1.21.1-0.20240913074525-1f5de7d620c4/go.mod h1:AFuwk3+uIWFeO5ohKjx5w606Djl940+wktaZ441VoCI= +github.com/livekit/protocol v1.21.1-0.20240918070421-da4a63076085 h1:dadfs8anNCNo2k4zWP2MJFASlXegQN5Sxw8aBDI0+0s= +github.com/livekit/protocol v1.21.1-0.20240918070421-da4a63076085/go.mod h1:AFuwk3+uIWFeO5ohKjx5w606Djl940+wktaZ441VoCI= github.com/livekit/psrpc v0.5.3-0.20240616012458-ac39c8549a0a h1:EQAHmcYEGlc6V517cQ3Iy0+jHgP6+tM/B4l2vGuLpQo= github.com/livekit/psrpc v0.5.3-0.20240616012458-ac39c8549a0a/go.mod h1:CQUBSPfYYAaevg1TNCc6/aYsa8DJH4jSRFdCeSZk5u0= github.com/mackerelio/go-osstat v0.2.5 h1:+MqTbZUhoIt4m8qzkVoXUJg1EuifwlAJSk4Yl2GXh+o= diff --git a/pkg/agent/agent_test.go b/pkg/agent/agent_test.go index 52dee7839..079cf8e02 100644 --- a/pkg/agent/agent_test.go +++ b/pkg/agent/agent_test.go @@ -8,9 +8,10 @@ import ( "time" "github.com/stretchr/testify/require" + "go.uber.org/atomic" "github.com/livekit/livekit-server/pkg/agent" - "github.com/livekit/livekit-server/pkg/agent/testutil" + "github.com/livekit/livekit-server/pkg/agent/testutils" "github.com/livekit/protocol/livekit" "github.com/livekit/protocol/rpc" "github.com/livekit/protocol/utils/guid" @@ -23,11 +24,11 @@ func TestAgent(t *testing.T) { bus := psrpc.NewLocalMessageBus() client := must.Get(rpc.NewAgentInternalClient(bus)) - server := testutil.NewTestServer(bus) + server := testutils.NewTestServer(bus) t.Cleanup(server.Close) worker := server.SimulateAgentWorker() - worker.Register("", "test", livekit.JobType_JT_ROOM) + worker.Register("test", livekit.JobType_JT_ROOM) jobAssignments := worker.JobAssignments.Observe() job := &livekit.Job{ @@ -35,7 +36,7 @@ func TestAgent(t *testing.T) { DispatchId: guid.New(guid.AgentDispatchPrefix), Type: livekit.JobType_JT_ROOM, Room: &livekit.Room{}, - Namespace: "test", + AgentName: "test", } _, err := client.JobRequest(context.Background(), "test", agent.RoomAgentTopic, job) require.NoError(t, err) @@ -51,7 +52,26 @@ func TestAgent(t *testing.T) { func TestAgentLoadBalancing(t *testing.T) { - batchJobCreate := func(wg *sync.WaitGroup, batchSize int, totalJobs int, client rpc.AgentInternalClient) { + batchJobCreate := func(batchSize int, totalJobs int, client rpc.AgentInternalClient, workers []*testutils.AgentWorker) <-chan struct{} { + var assigned atomic.Uint32 + done := make(chan struct{}) + for _, w := range workers { + assignments := w.JobAssignments.Observe() + go func() { + defer assignments.Stop() + for { + select { + case <-done: + case <-assignments.Events(): + if assigned.Inc() == uint32(totalJobs) { + close(done) + } + } + } + }() + } + + var wg sync.WaitGroup for i := 0; i < totalJobs; i += batchSize { wg.Add(1) go func(start int) { @@ -62,13 +82,16 @@ func TestAgentLoadBalancing(t *testing.T) { DispatchId: guid.New(guid.AgentDispatchPrefix), Type: livekit.JobType_JT_ROOM, Room: &livekit.Room{}, - Namespace: "test", + AgentName: "test", } _, err := client.JobRequest(context.Background(), "test", agent.RoomAgentTopic, job) require.NoError(t, err) } }(i) } + wg.Wait() + + return done } t.Run("jobs are distributed normally with baseline worker load", func(t *testing.T) { @@ -78,50 +101,35 @@ func TestAgentLoadBalancing(t *testing.T) { bus := psrpc.NewLocalMessageBus() client := must.Get(rpc.NewAgentInternalClient(bus)) - server := testutil.NewTestServer(bus) + server := testutils.NewTestServer(bus) t.Cleanup(server.Close) - agents := make([]*testutil.AgentWorker, totalWorkers) + agents := make([]*testutils.AgentWorker, totalWorkers) for i := 0; i < totalWorkers; i++ { - agents[i] = server.SimulateAgentWorker() - agents[i].Register(fmt.Sprintf("agent-%d", i), "test", livekit.JobType_JT_ROOM) + agents[i] = server.SimulateAgentWorker( + testutils.WithLabel(fmt.Sprintf("agent-%d", i)), + testutils.WithJobLoad(testutils.NewStableJobLoad(0.01)), + ) + agents[i].Register("test", livekit.JobType_JT_ROOM) } - jobAssignments := make(chan *livekit.Job, totalJobs) - for i := 0; i < totalWorkers; i++ { - worker := agents[i] - go func() { - for a := range worker.JobAssignments.Observe().Events() { - jobAssignments <- a.Job - } - }() + select { + case <-batchJobCreate(10, totalJobs, client, agents): + case <-time.After(time.Second): + require.Fail(t, "job assignment timeout") } - var wg sync.WaitGroup - batchJobCreate(&wg, 10, totalJobs, client) - wg.Wait() - jobCount := make(map[string]int) - for i := 0; i < totalJobs; i++ { - select { - case job := <-jobAssignments: - jobCount[job.AgentName]++ - case <-time.After(time.Second): - require.Fail(t, "job assignment timeout") - } + for _, w := range agents { + jobCount[w.Label] = len(w.Jobs()) } - assignedJobs := 0 // check that jobs are distributed normally for i := 0; i < totalWorkers; i++ { - agentName := fmt.Sprintf("agent-%d", i) - assignedJobs += jobCount[agentName] - require.GreaterOrEqual(t, jobCount[agentName], 0) - require.Less(t, jobCount[agentName], 35) // three std deviations from the mean is 32 + label := fmt.Sprintf("agent-%d", i) + require.GreaterOrEqual(t, jobCount[label], 0) + require.Less(t, jobCount[label], 35) // three std deviations from the mean is 32 } - - // ensure all jobs are assigned - require.Equal(t, 100, assignedJobs) }) t.Run("jobs are distributed with variable and overloaded worker load", func(t *testing.T) { @@ -131,58 +139,41 @@ func TestAgentLoadBalancing(t *testing.T) { bus := psrpc.NewLocalMessageBus() client := must.Get(rpc.NewAgentInternalClient(bus)) - server := testutil.NewTestServer(bus) + server := testutils.NewTestServer(bus) t.Cleanup(server.Close) - agents := make([]*testutil.AgentWorker, totalWorkers) + agents := make([]*testutils.AgentWorker, totalWorkers) for i := 0; i < totalWorkers; i++ { + label := fmt.Sprintf("agent-%d", i) if i%2 == 0 { // make sure we have some workers that can accept jobs - agents[i] = server.SimulateAgentWorker() + agents[i] = server.SimulateAgentWorker(testutils.WithLabel(label)) } else { - agents[i] = server.SimulateAgentWorker(testutil.WithDefaultWorkerLoad(0.9)) + agents[i] = server.SimulateAgentWorker(testutils.WithLabel(label), testutils.WithDefaultWorkerLoad(0.9)) } - agents[i].Register(fmt.Sprintf("agent-%d", i), "test", livekit.JobType_JT_ROOM) + agents[i].Register("test", livekit.JobType_JT_ROOM) } - jobAssignments := make(chan *livekit.Job, totalJobs) - for i := 0; i < totalWorkers; i++ { - worker := agents[i] - go func() { - for a := range worker.JobAssignments.Observe().Events() { - jobAssignments <- a.Job - } - }() + select { + case <-batchJobCreate(1, totalJobs, client, agents): + case <-time.After(time.Second): + require.Fail(t, "job assignment timeout") } - var wg sync.WaitGroup - batchJobCreate(&wg, 1, totalJobs, client) - wg.Wait() - jobCount := make(map[string]int) - for i := 0; i < totalJobs; i++ { - select { - case job := <-jobAssignments: - jobCount[job.AgentName]++ - case <-time.After(time.Second): - require.Fail(t, "job assignment timeout") - } + for _, w := range agents { + jobCount[w.Label] = len(w.Jobs()) } - assignedJobs := 0 for i := 0; i < totalWorkers; i++ { - agentName := fmt.Sprintf("agent-%d", i) - assignedJobs += jobCount[agentName] + label := fmt.Sprintf("agent-%d", i) if i%2 == 0 { - require.GreaterOrEqual(t, jobCount[agentName], 2) + require.GreaterOrEqual(t, jobCount[label], 2) } else { - require.Equal(t, 0, jobCount[agentName]) + require.Equal(t, 0, jobCount[label]) } - require.GreaterOrEqual(t, jobCount[agentName], 0) + require.GreaterOrEqual(t, jobCount[label], 0) } - - // ensure all jobs are assigned - require.Equal(t, 15, assignedJobs) }) } diff --git a/pkg/agent/testutil/server.go b/pkg/agent/testutils/server.go similarity index 88% rename from pkg/agent/testutil/server.go rename to pkg/agent/testutils/server.go index 04f437534..4178e0100 100644 --- a/pkg/agent/testutil/server.go +++ b/pkg/agent/testutils/server.go @@ -1,4 +1,4 @@ -package testutil +package testutils import ( "context" @@ -11,41 +11,45 @@ import ( "github.com/frostbyte73/core" "github.com/gammazero/deque" + "golang.org/x/exp/maps" "github.com/livekit/livekit-server/pkg/agent" "github.com/livekit/livekit-server/pkg/config" "github.com/livekit/livekit-server/pkg/service" "github.com/livekit/protocol/auth" "github.com/livekit/protocol/livekit" - "github.com/livekit/protocol/logger" "github.com/livekit/protocol/utils" "github.com/livekit/protocol/utils/guid" "github.com/livekit/protocol/utils/must" + "github.com/livekit/protocol/utils/options" "github.com/livekit/psrpc" ) +type AgentService interface { + HandleConnection(context.Context, agent.SignalConn, agent.WorkerProtocolVersion) + DrainConnections(time.Duration) +} + type TestServer struct { - *service.AgentService - keyProvider auth.KeyProvider + AgentService } func NewTestServer(bus psrpc.MessageBus) *TestServer { - keyProvider := auth.NewSimpleKeyProvider("test", "verysecretsecret") - - s := must.Get(service.NewAgentService( + return NewTestServerWithService(must.Get(service.NewAgentService( &config.Config{Region: "test"}, &livekit.Node{Id: guid.New("N_")}, bus, - keyProvider, - )) + auth.NewSimpleKeyProvider("test", "verysecretsecret"), + ))) +} - return &TestServer{ - AgentService: s, - keyProvider: keyProvider, - } +func NewTestServerWithService(s AgentService) *TestServer { + return &TestServer{s} } type SimulatedWorkerOptions struct { + Context context.Context + Label string SupportResume bool DefaultJobLoad float32 JobLoadThreshold float32 @@ -56,6 +60,18 @@ type SimulatedWorkerOptions struct { type SimulatedWorkerOption func(*SimulatedWorkerOptions) +func WithContext(ctx context.Context) SimulatedWorkerOption { + return func(o *SimulatedWorkerOptions) { + o.Context = ctx + } +} + +func WithLabel(label string) SimulatedWorkerOption { + return func(o *SimulatedWorkerOptions) { + o.Label = label + } +} + func WithJobAvailabilityHandler(h func(AgentJobRequest)) SimulatedWorkerOption { return func(o *SimulatedWorkerOptions) { o.HandleAvailability = h @@ -80,15 +96,15 @@ func WithDefaultWorkerLoad(load float32) SimulatedWorkerOption { func (h *TestServer) SimulateAgentWorker(opts ...SimulatedWorkerOption) *AgentWorker { o := &SimulatedWorkerOptions{ + Context: context.Background(), + Label: guid.New("TEST_AGENT_"), DefaultJobLoad: 0.1, JobLoadThreshold: 0.8, DefaultWorkerLoad: 0.0, HandleAvailability: func(r AgentJobRequest) { r.Accept() }, HandleAssignment: func(j *livekit.Job) JobLoad { return nil }, } - for _, opt := range opts { - opt(o) - } + options.Apply(o, opts) w := &AgentWorker{ workerMessages: make(chan *livekit.WorkerMessage, 1), @@ -107,47 +123,14 @@ func (h *TestServer) SimulateAgentWorker(opts ...SimulatedWorkerOption) *AgentWo w.sendStatus() } - go w.worker() - go h.handleConnection(w) + ctx := service.WithAPIKey(o.Context, &auth.ClaimGrants{}, "test") + go h.HandleConnection(ctx, w, agent.CurrentProtocol) + return w } -func (h *TestServer) handleConnection(w *AgentWorker) { - worker := agent.NewWorker( - agent.CurrentProtocol, - "test", - h.keyProvider.GetSecret("test"), - &livekit.ServerInfo{}, - w, - logger.GetLogger(), - h, - ) - - h.InsertWorker(worker) - - for { - req, _, err := w.ReadWorkerMessage() - if err != nil { - if service.IsWebSocketCloseError(err) { - worker.Logger().Infow("worker closed WS connection", "wsError", err) - } else { - worker.Logger().Errorw("error reading from websocket", err) - } - break - } - - worker.HandleMessage(req) - } - - h.DeleteWorker(worker) - - worker.Close() -} - func (h *TestServer) Close() { - for _, w := range h.Workers() { - w.Close() - } + h.DrainConnections(1) } var _ agent.SignalConn = (*AgentWorker)(nil) @@ -185,7 +168,6 @@ func (r AgentJobRequest) Reject() { } type AgentWorker struct { - Name string *SimulatedWorkerOptions fuse core.Fuse @@ -203,13 +185,13 @@ type AgentWorker struct { WorkerPongs *utils.EventObserverList[*livekit.WorkerPong] } -func (w *AgentWorker) worker() { - t := time.NewTicker(5 * time.Second) +func (w *AgentWorker) statusWorker() { + t := time.NewTicker(2 * time.Second) defer t.Stop() for !w.fuse.IsBroken() { - <-t.C w.sendStatus() + <-t.C } } @@ -303,7 +285,6 @@ func (w *AgentWorker) handleAvailability(m *livekit.AvailabilityRequest) { } func (w *AgentWorker) handleAssignment(m *livekit.JobAssignment) { - m.Job.AgentName = w.Name w.JobAssignments.Emit(m) var load JobLoad @@ -407,13 +388,12 @@ func (w *AgentWorker) sendStatus() { }) } -func (w *AgentWorker) Register(name string, namespace string, jobType livekit.JobType) { - w.Name = name +func (w *AgentWorker) Register(agentName string, jobType livekit.JobType) { w.SendRegister(&livekit.RegisterWorkerRequest{ Type: jobType, - Namespace: &namespace, + AgentName: agentName, }) - w.sendStatus() + go w.statusWorker() } func (w *AgentWorker) SimulateRoomJob(roomName string) { @@ -426,6 +406,12 @@ func (w *AgentWorker) SimulateRoomJob(roomName string) { }) } +func (w *AgentWorker) Jobs() []*AgentJob { + w.mu.Lock() + defer w.mu.Unlock() + return maps.Values(w.jobs) +} + type stableJobLoad struct { load float32 } @@ -490,6 +476,6 @@ func NewNormalRandomJobLoadWithRNG(mean, stddev float64, rng *rand.Rand) JobLoad func (s normalRandomJobLoad) Load() float32 { u := 1 - s.rng() v := s.rng() - z := math.Sqrt(-2.0*math.Log(u)) * math.Cos(2.0*math.Pi*v) + z := math.Sqrt(-2*math.Log(u)) * math.Cos(2*math.Pi*v) return float32(max(0, z*s.stddev+s.mean)) } diff --git a/pkg/agent/worker.go b/pkg/agent/worker.go index 7f3f201ae..7de167374 100644 --- a/pkg/agent/worker.go +++ b/pkg/agent/worker.go @@ -17,10 +17,12 @@ package agent import ( "context" "errors" + "fmt" "sync" - "sync/atomic" "time" + "google.golang.org/protobuf/proto" + pagent "github.com/livekit/protocol/agent" "github.com/livekit/protocol/livekit" "github.com/livekit/protocol/logger" @@ -29,6 +31,16 @@ import ( "github.com/livekit/psrpc" ) +var ( + ErrUnimplementedWrorkerSignal = errors.New("unimplemented worker signal") + ErrUnknownWorkerSignal = errors.New("unknown worker signal") + ErrUnknownJobType = errors.New("unknown job type") + ErrWorkerClosed = errors.New("worker closed") + ErrWorkerNotAvailable = errors.New("worker not available") + ErrAvailabilityTimeout = errors.New("agent worker availability timeout") + ErrDuplicateJobAssignment = errors.New("duplicate job assignment") +) + type WorkerProtocolVersion int const CurrentProtocol = 1 @@ -39,110 +51,221 @@ const ( pingFrequency = 10 * time.Second ) -var ( - ErrWorkerClosed = errors.New("worker closed") - ErrWorkerNotAvailable = errors.New("worker not available") - ErrAvailabilityTimeout = errors.New("agent worker availability timeout") - ErrDuplicateJobAssignment = errors.New("duplicate job assignment") -) - type SignalConn interface { WriteServerMessage(msg *livekit.ServerMessage) (int, error) ReadWorkerMessage() (*livekit.WorkerMessage, int, error) + SetReadDeadline(time.Time) error Close() error } -type WorkerHandler interface { - HandleWorkerRegister(w *Worker) - HandleWorkerDeregister(w *Worker) - HandleWorkerStatus(w *Worker, status *livekit.UpdateWorkerStatus) - HandleWorkerJobStatus(w *Worker, status *livekit.UpdateJobStatus) - HandleWorkerSimulateJob(w *Worker, job *livekit.Job) - HandleWorkerMigrateJob(w *Worker, request *livekit.MigrateJobRequest) -} - -var _ WorkerHandler = UnimplementedWorkerHandler{} - -type UnimplementedWorkerHandler struct{} - -func (UnimplementedWorkerHandler) HandleWorkerRegister(*Worker) {} -func (UnimplementedWorkerHandler) HandleWorkerDeregister(*Worker) {} -func (UnimplementedWorkerHandler) HandleWorkerStatus(*Worker, *livekit.UpdateWorkerStatus) {} -func (UnimplementedWorkerHandler) HandleWorkerJobStatus(*Worker, *livekit.UpdateJobStatus) {} -func (UnimplementedWorkerHandler) HandleWorkerSimulateJob(*Worker, *livekit.Job) {} -func (UnimplementedWorkerHandler) HandleWorkerMigrateJob(*Worker, *livekit.MigrateJobRequest) {} - func JobStatusIsEnded(s livekit.JobStatus) bool { return s == livekit.JobStatus_JS_SUCCESS || s == livekit.JobStatus_JS_FAILED } +type WorkerSignalHandler interface { + HandleRegister(*livekit.RegisterWorkerRequest) error + HandleAvailability(*livekit.AvailabilityResponse) error + HandleUpdateJob(*livekit.UpdateJobStatus) error + HandleSimulateJob(*livekit.SimulateJobRequest) error + HandlePing(*livekit.WorkerPing) error + HandleUpdateWorker(*livekit.UpdateWorkerStatus) error + HandleMigrateJob(*livekit.MigrateJobRequest) error +} + +func DispatchWorkerSignal(req *livekit.WorkerMessage, h WorkerSignalHandler) error { + switch m := req.Message.(type) { + case *livekit.WorkerMessage_Register: + return h.HandleRegister(m.Register) + case *livekit.WorkerMessage_Availability: + return h.HandleAvailability(m.Availability) + case *livekit.WorkerMessage_UpdateJob: + return h.HandleUpdateJob(m.UpdateJob) + case *livekit.WorkerMessage_SimulateJob: + return h.HandleSimulateJob(m.SimulateJob) + case *livekit.WorkerMessage_Ping: + return h.HandlePing(m.Ping) + case *livekit.WorkerMessage_UpdateWorker: + return h.HandleUpdateWorker(m.UpdateWorker) + case *livekit.WorkerMessage_MigrateJob: + return h.HandleMigrateJob(m.MigrateJob) + default: + return ErrUnknownWorkerSignal + } +} + +var _ WorkerSignalHandler = (*UnimplementedWorkerSignalHandler)(nil) + +type UnimplementedWorkerSignalHandler struct{} + +func (UnimplementedWorkerSignalHandler) HandleRegister(*livekit.RegisterWorkerRequest) error { + return fmt.Errorf("%w: Register", ErrUnimplementedWrorkerSignal) +} +func (UnimplementedWorkerSignalHandler) HandleAvailability(*livekit.AvailabilityResponse) error { + return fmt.Errorf("%w: Availability", ErrUnimplementedWrorkerSignal) +} +func (UnimplementedWorkerSignalHandler) HandleUpdateJob(*livekit.UpdateJobStatus) error { + return fmt.Errorf("%w: UpdateJob", ErrUnimplementedWrorkerSignal) +} +func (UnimplementedWorkerSignalHandler) HandleSimulateJob(*livekit.SimulateJobRequest) error { + return fmt.Errorf("%w: SimulateJob", ErrUnimplementedWrorkerSignal) +} +func (UnimplementedWorkerSignalHandler) HandlePing(*livekit.WorkerPing) error { + return fmt.Errorf("%w: Ping", ErrUnimplementedWrorkerSignal) +} +func (UnimplementedWorkerSignalHandler) HandleUpdateWorker(*livekit.UpdateWorkerStatus) error { + return fmt.Errorf("%w: UpdateWorker", ErrUnimplementedWrorkerSignal) +} +func (UnimplementedWorkerSignalHandler) HandleMigrateJob(*livekit.MigrateJobRequest) error { + return fmt.Errorf("%w: MigrateJob", ErrUnimplementedWrorkerSignal) +} + +type WorkerPingHandler struct { + UnimplementedWorkerSignalHandler + conn SignalConn +} + +func (h WorkerPingHandler) HandlePing(ping *livekit.WorkerPing) error { + _, err := h.conn.WriteServerMessage(&livekit.ServerMessage{ + Message: &livekit.ServerMessage_Pong{ + Pong: &livekit.WorkerPong{ + LastTimestamp: ping.Timestamp, + Timestamp: time.Now().UnixMilli(), + }, + }, + }) + return err +} + +type WorkerRegistration struct { + Protocol WorkerProtocolVersion + ID string + Version string + AgentName string + Namespace string + JobType livekit.JobType + Permissions *livekit.ParticipantPermission +} + +var _ WorkerSignalHandler = (*WorkerRegisterer)(nil) + +type WorkerRegisterer struct { + WorkerPingHandler + serverInfo *livekit.ServerInfo + protocol WorkerProtocolVersion + deadline time.Time + + registration WorkerRegistration + registered bool +} + +func NewWorkerRegisterer(conn SignalConn, serverInfo *livekit.ServerInfo, protocol WorkerProtocolVersion) *WorkerRegisterer { + return &WorkerRegisterer{ + WorkerPingHandler: WorkerPingHandler{conn: conn}, + serverInfo: serverInfo, + protocol: protocol, + deadline: time.Now().Add(registerTimeout), + } +} + +func (h *WorkerRegisterer) Deadline() time.Time { + return h.deadline +} + +func (h *WorkerRegisterer) Registration() WorkerRegistration { + return h.registration +} + +func (h *WorkerRegisterer) Registered() bool { + return h.registered +} + +func (h *WorkerRegisterer) HandleRegister(req *livekit.RegisterWorkerRequest) error { + if !livekit.IsJobType(req.GetType()) { + return ErrUnknownJobType + } + + permissions := req.AllowedPermissions + if permissions == nil { + permissions = &livekit.ParticipantPermission{ + CanSubscribe: true, + CanPublish: true, + CanPublishData: true, + CanUpdateMetadata: true, + } + } + + h.registration = WorkerRegistration{ + Protocol: h.protocol, + ID: guid.New(guid.AgentWorkerPrefix), + Version: req.Version, + AgentName: req.AgentName, + Namespace: req.GetNamespace(), + JobType: req.GetType(), + Permissions: permissions, + } + h.registered = true + + _, err := h.conn.WriteServerMessage(&livekit.ServerMessage{ + Message: &livekit.ServerMessage_Register{ + Register: &livekit.RegisterWorkerResponse{ + WorkerId: h.registration.ID, + ServerInfo: h.serverInfo, + }, + }, + }) + return err +} + +var _ WorkerSignalHandler = (*Worker)(nil) + type Worker struct { - id string - jobType livekit.JobType - version string - agentName string - namespace string - load float32 - permissions *livekit.ParticipantPermission - apiKey string - apiSecret string - serverInfo *livekit.ServerInfo - mu sync.Mutex + WorkerPingHandler + WorkerRegistration - protocolVersion WorkerProtocolVersion - registered atomic.Bool - status livekit.WorkerStatus - runningJobs map[string]*livekit.Job // JobID -> Job - - handler WorkerHandler - - conn SignalConn - closed chan struct{} - - availability map[string]chan *livekit.AvailabilityResponse + apiKey string + apiSecret string + conn SignalConn + logger logger.Logger ctx context.Context cancel context.CancelFunc + closed chan struct{} - logger logger.Logger + mu sync.Mutex + load float32 + status livekit.WorkerStatus + + runningJobs map[string]*livekit.Job + availability map[string]chan *livekit.AvailabilityResponse } func NewWorker( - protocolVersion WorkerProtocolVersion, + registration WorkerRegistration, apiKey string, apiSecret string, - serverInfo *livekit.ServerInfo, conn SignalConn, logger logger.Logger, - handler WorkerHandler, ) *Worker { ctx, cancel := context.WithCancel(context.Background()) - id := guid.New(guid.AgentWorkerPrefix) - w := &Worker{ - id: id, - protocolVersion: protocolVersion, - apiKey: apiKey, - apiSecret: apiSecret, - serverInfo: serverInfo, - closed: make(chan struct{}), - runningJobs: make(map[string]*livekit.Job), - availability: make(map[string]chan *livekit.AvailabilityResponse), - conn: conn, - ctx: ctx, - cancel: cancel, - logger: logger.WithValues("workerID", id), - handler: handler, + return &Worker{ + WorkerRegistration: registration, + apiKey: apiKey, + apiSecret: apiSecret, + conn: conn, + logger: logger.WithValues( + "workerID", registration.ID, + "agentName", registration.AgentName, + "jobType", registration.JobType.String(), + ), + + ctx: ctx, + cancel: cancel, + closed: make(chan struct{}), + + runningJobs: make(map[string]*livekit.Job), + availability: make(map[string]chan *livekit.AvailabilityResponse), } - - time.AfterFunc(registerTimeout, func() { - if !w.registered.Load() && !w.IsClosed() { - w.logger.Warnw("worker did not register in time", nil, "id", w.id) - w.Close() - } - }) - - return w } func (w *Worker) sendRequest(req *livekit.ServerMessage) { @@ -151,24 +274,6 @@ func (w *Worker) sendRequest(req *livekit.ServerMessage) { } } -func (w *Worker) ID() string { - return w.id -} - -func (w *Worker) JobType() livekit.JobType { - return w.jobType -} - -func (w *Worker) Namespace() string { - return w.namespace -} - -func (w *Worker) AgentName() string { - w.mu.Lock() - defer w.mu.Unlock() - return w.agentName -} - func (w *Worker) Status() livekit.WorkerStatus { w.mu.Lock() defer w.mu.Unlock() @@ -236,7 +341,7 @@ func (w *Worker) AssignJob(ctx context.Context, job *livekit.Job) error { job.State.ParticipantIdentity = res.ParticipantIdentity - token, err := pagent.BuildAgentToken(w.apiKey, w.apiSecret, job.Room.Name, res.ParticipantIdentity, res.ParticipantName, res.ParticipantMetadata, w.permissions) + token, err := pagent.BuildAgentToken(w.apiKey, w.apiSecret, job.Room.Name, res.ParticipantIdentity, res.ParticipantName, res.ParticipantMetadata, w.Permissions) if err != nil { w.logger.Errorw("failed to build agent token", err) return err @@ -285,13 +390,11 @@ func (w *Worker) TerminateJob(jobID string, reason rpc.JobTerminateReason) (*liv errorStr = "agent worker left the room" } - w.updateJobStatus(&livekit.UpdateJobStatus{ + return w.UpdateJobStatus(&livekit.UpdateJobStatus{ JobId: jobID, Status: status, Error: errorStr, }) - - return job.State, nil } func (w *Worker) UpdateMetadata(metadata string) { @@ -320,116 +423,46 @@ func (w *Worker) Close() { w.cancel() _ = w.conn.Close() w.mu.Unlock() - - if w.registered.Load() { - w.handler.HandleWorkerDeregister(w) - } } -func (w *Worker) HandleMessage(req *livekit.WorkerMessage) { - switch m := req.Message.(type) { - case *livekit.WorkerMessage_Register: - w.handleRegister(m.Register) - case *livekit.WorkerMessage_Availability: - w.handleAvailability(m.Availability) - case *livekit.WorkerMessage_UpdateJob: - w.handleJobUpdate(m.UpdateJob) - case *livekit.WorkerMessage_SimulateJob: - w.handleSimulateJob(m.SimulateJob) - case *livekit.WorkerMessage_Ping: - w.handleWorkerPing(m.Ping) - case *livekit.WorkerMessage_UpdateWorker: - w.handleWorkerStatus(m.UpdateWorker) - case *livekit.WorkerMessage_MigrateJob: - w.handleMigrateJob(m.MigrateJob) - } -} - -func (w *Worker) handleRegister(req *livekit.RegisterWorkerRequest) { - w.mu.Lock() - var err error - if w.IsClosed() { - err = errors.New("worker closed") - } - if w.registered.Swap(true) { - err = errors.New("worker already registered") - } - if err != nil { - w.mu.Unlock() - w.logger.Warnw("unable to register worker", err, "id", w.id) - return - } - - w.version = req.Version - w.agentName = req.GetAgentName() - w.namespace = req.GetNamespace() - w.jobType = req.GetType() - - if req.AllowedPermissions != nil { - w.permissions = req.AllowedPermissions - } else { - // Use default agent permissions - w.permissions = &livekit.ParticipantPermission{ - CanSubscribe: true, - CanPublish: true, - CanPublishData: true, - CanUpdateMetadata: true, - } - } - - w.status = livekit.WorkerStatus_WS_AVAILABLE - w.mu.Unlock() - - w.logger.Debugw("worker registered", "request", logger.Proto(req)) - - w.sendRequest(&livekit.ServerMessage{ - Message: &livekit.ServerMessage_Register{ - Register: &livekit.RegisterWorkerResponse{ - WorkerId: w.ID(), - ServerInfo: w.serverInfo, - }, - }, - }) - - w.handler.HandleWorkerRegister(w) -} - -func (w *Worker) handleAvailability(res *livekit.AvailabilityResponse) { +func (w *Worker) HandleAvailability(res *livekit.AvailabilityResponse) error { w.mu.Lock() defer w.mu.Unlock() availCh, ok := w.availability[res.JobId] if !ok { w.logger.Warnw("received availability response for unknown job", nil, "jobId", res.JobId) - return + return nil } availCh <- res delete(w.availability, res.JobId) + + return nil } -func (w *Worker) handleJobUpdate(update *livekit.UpdateJobStatus) { - err := w.updateJobStatus(update) +func (w *Worker) HandleUpdateJob(update *livekit.UpdateJobStatus) error { + _, err := w.UpdateJobStatus(update) if err != nil { w.logger.Infow("received job update for unknown job", "jobID", update.JobId) } + return err } -func (w *Worker) updateJobStatus(update *livekit.UpdateJobStatus) error { +func (w *Worker) UpdateJobStatus(update *livekit.UpdateJobStatus) (*livekit.JobState, error) { w.mu.Lock() + defer w.mu.Unlock() job, ok := w.runningJobs[update.JobId] if !ok { - w.mu.Unlock() - - return psrpc.NewErrorf(psrpc.NotFound, "received job update for unknown job") + return nil, psrpc.NewErrorf(psrpc.NotFound, "received job update for unknown job") } now := time.Now() job.State.UpdatedAt = now.UnixNano() - if job.State.Status == livekit.JobStatus_JS_PENDING && JobStatusIsEnded(update.Status) { + if job.State.Status == livekit.JobStatus_JS_PENDING && update.Status != livekit.JobStatus_JS_PENDING { job.State.StartedAt = now.UnixNano() } @@ -444,14 +477,11 @@ func (w *Worker) updateJobStatus(update *livekit.UpdateJobStatus) error { if JobStatusIsEnded(job.State.Status) { delete(w.runningJobs, job.Id) } - w.mu.Unlock() - w.handler.HandleWorkerJobStatus(w, update) - - return nil + return proto.Clone(job.State).(*livekit.JobState), nil } -func (w *Worker) handleSimulateJob(simulate *livekit.SimulateJobRequest) { +func (w *Worker) HandleSimulateJob(simulate *livekit.SimulateJobRequest) error { jobType := livekit.JobType_JT_ROOM if simulate.Participant != nil { jobType = livekit.JobType_JT_PUBLISHER @@ -462,44 +492,36 @@ func (w *Worker) handleSimulateJob(simulate *livekit.SimulateJobRequest) { Type: jobType, Room: simulate.Room, Participant: simulate.Participant, - Namespace: w.Namespace(), - AgentName: w.AgentName(), + Namespace: w.Namespace, + AgentName: w.AgentName, } go func() { err := w.AssignJob(w.ctx, job) if err != nil { w.logger.Errorw("failed to simulate job, assignment failed", err, "jobId", job.Id) - } else { - w.handler.HandleWorkerSimulateJob(w, job) } }() + + return nil } -func (w *Worker) handleWorkerPing(ping *livekit.WorkerPing) { - w.sendRequest(&livekit.ServerMessage{Message: &livekit.ServerMessage_Pong{ - Pong: &livekit.WorkerPong{ - LastTimestamp: ping.Timestamp, - Timestamp: time.Now().UnixMilli(), - }, - }}) -} - -func (w *Worker) handleWorkerStatus(update *livekit.UpdateWorkerStatus) { +func (w *Worker) HandleUpdateWorker(update *livekit.UpdateWorkerStatus) error { w.logger.Debugw("worker status update", "update", logger.Proto(update)) w.mu.Lock() + defer w.mu.Unlock() + if update.Status != nil { w.status = update.GetStatus() } w.load = update.GetLoad() - w.mu.Unlock() - w.handler.HandleWorkerStatus(w, update) + return nil } -func (w *Worker) handleMigrateJob(migrate *livekit.MigrateJobRequest) { +func (w *Worker) HandleMigrateJob(req *livekit.MigrateJobRequest) error { // TODO(theomonnom): On OSS this is not implemented // We could maybe just move a specific job to another worker - w.handler.HandleWorkerMigrateJob(w, migrate) + return nil } diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index 7aacdd2f3..e7e495994 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -39,6 +39,7 @@ type WebsocketClient interface { ReadMessage() (messageType int, p []byte, err error) WriteMessage(messageType int, data []byte) error WriteControl(messageType int, data []byte, deadline time.Time) error + SetReadDeadline(deadline time.Time) error Close() error } diff --git a/pkg/service/agentservice.go b/pkg/service/agentservice.go index 4762731b7..f6491230a 100644 --- a/pkg/service/agentservice.go +++ b/pkg/service/agentservice.go @@ -26,7 +26,6 @@ import ( "time" "github.com/gorilla/websocket" - "golang.org/x/exp/maps" "google.golang.org/protobuf/types/known/emptypb" "github.com/livekit/livekit-server/pkg/agent" @@ -49,6 +48,14 @@ type AgentSocketUpgrader struct { } func (u AgentSocketUpgrader) Upgrade(w http.ResponseWriter, r *http.Request, responseHeader http.Header) (*websocket.Conn, agent.WorkerProtocolVersion, bool) { + if u.CheckOrigin == nil { + // allow connections from any origin, since script may be hosted anywhere + // security is enforced by access tokens + u.CheckOrigin = func(r *http.Request) bool { + return true + } + } + // reject non websocket requests if !websocket.IsWebSocketUpgrade(r) { w.WriteHeader(404) @@ -77,6 +84,41 @@ func (u AgentSocketUpgrader) Upgrade(w http.ResponseWriter, r *http.Request, res return conn, protocol, true } +func DispatchAgentWorkerSignal(c agent.SignalConn, h agent.WorkerSignalHandler, l logger.Logger) bool { + req, _, err := c.ReadWorkerMessage() + if err != nil { + if IsWebSocketCloseError(err) { + l.Infow("worker closed WS connection", "wsError", err) + } else { + l.Errorw("error reading from websocket", err) + } + return false + } + + if err := agent.DispatchWorkerSignal(req, h); err != nil { + l.Warnw("unable to handle worker signal", err, "req", logger.Proto(req)) + return false + } + + return true +} + +func HandshakeAgentWorker(c agent.SignalConn, serverInfo *livekit.ServerInfo, protocol agent.WorkerProtocolVersion, l logger.Logger) (r agent.WorkerRegistration, ok bool) { + wr := agent.NewWorkerRegisterer(c, serverInfo, protocol) + if err := c.SetReadDeadline(wr.Deadline()); err != nil { + return + } + for !wr.Registered() { + if ok = DispatchAgentWorkerSignal(c, wr, l); !ok { + return + } + } + if err := c.SetReadDeadline(time.Time{}); err != nil { + return + } + return wr.Registration(), true +} + type AgentService struct { upgrader AgentSocketUpgrader @@ -84,8 +126,6 @@ type AgentService struct { } type AgentHandler struct { - agent.UnimplementedWorkerHandler - agentServer rpc.AgentInternalServer mu sync.Mutex logger logger.Logger @@ -118,12 +158,6 @@ func NewAgentService(conf *config.Config, ) (*AgentService, error) { s := &AgentService{} - // allow connections from any origin, since script may be hosted anywhere - // security is enforced by access tokens - s.upgrader.CheckOrigin = func(r *http.Request) bool { - return true - } - serverInfo := &livekit.ServerInfo{ Edition: livekit.ServerInfo_Standard, Version: version.Version, @@ -151,6 +185,7 @@ func NewAgentService(conf *config.Config, func (s *AgentService) ServeHTTP(writer http.ResponseWriter, r *http.Request) { if conn, protocol, ok := s.upgrader.Upgrade(writer, r, nil); ok { s.HandleConnection(r.Context(), NewWSSignalConnection(conn), protocol) + conn.Close() } } @@ -175,63 +210,40 @@ func NewAgentHandler( } } -func (h *AgentHandler) InsertWorker(w *agent.Worker) { - h.mu.Lock() - defer h.mu.Unlock() - h.workers[w.ID()] = w -} - -func (h *AgentHandler) DeleteWorker(w *agent.Worker) { - h.mu.Lock() - defer h.mu.Unlock() - delete(h.workers, w.ID()) -} - -func (h *AgentHandler) Workers() []*agent.Worker { - h.mu.Lock() - defer h.mu.Unlock() - return maps.Values(h.workers) -} - func (h *AgentHandler) HandleConnection(ctx context.Context, conn agent.SignalConn, protocol agent.WorkerProtocolVersion) { + registration, ok := HandshakeAgentWorker(conn, h.serverInfo, protocol, h.logger) + if !ok { + return + } + apiKey := GetAPIKey(ctx) apiSecret := h.keyProvider.GetSecret(apiKey) - worker := agent.NewWorker(protocol, apiKey, apiSecret, h.serverInfo, conn, h.logger, h) + worker := agent.NewWorker(registration, apiKey, apiSecret, conn, h.logger) + h.registerWorker(worker) - h.InsertWorker(worker) - - for { - req, _, err := conn.ReadWorkerMessage() - if err != nil { - if IsWebSocketCloseError(err) { - worker.Logger().Infow("worker closed WS connection", "wsError", err) - } else { - worker.Logger().Errorw("error reading from websocket", err) - } - break - } - - worker.HandleMessage(req) + for ok := true; ok; { + ok = DispatchAgentWorkerSignal(conn, worker, worker.Logger()) } - h.DeleteWorker(worker) - + h.deregisterWorker(worker) worker.Close() } -func (h *AgentHandler) HandleWorkerRegister(w *agent.Worker) { +func (h *AgentHandler) registerWorker(w *agent.Worker) { h.mu.Lock() - key := workerKey{w.AgentName(), w.Namespace(), w.JobType()} + h.workers[w.ID] = w + + key := workerKey{w.AgentName, w.Namespace, w.JobType} workers := h.namespaceWorkers[key] created := len(workers) == 0 if created { - nameTopic := agent.GetAgentTopic(w.AgentName(), w.Namespace()) + nameTopic := agent.GetAgentTopic(w.AgentName, w.Namespace) typeTopic := h.roomTopic - if w.JobType() == livekit.JobType_JT_PUBLISHER { + if w.JobType == livekit.JobType_JT_PUBLISHER { typeTopic = h.publisherTopic } err := h.agentServer.RegisterJobRequestTopic(nameTopic, typeTopic) @@ -243,36 +255,37 @@ func (h *AgentHandler) HandleWorkerRegister(w *agent.Worker) { return } - if w.JobType() == livekit.JobType_JT_ROOM { + if w.JobType == livekit.JobType_JT_ROOM { h.roomKeyCount++ } else { h.publisherKeyCount++ } - h.namespaces = append(h.namespaces, w.Namespace()) + h.namespaces = append(h.namespaces, w.Namespace) sort.Strings(h.namespaces) - h.agentNames = append(h.agentNames, w.AgentName()) + h.agentNames = append(h.agentNames, w.AgentName) sort.Strings(h.agentNames) - } h.namespaceWorkers[key] = append(workers, w) h.mu.Unlock() if created { - h.logger.Infow("initial worker registered", "namespace", w.Namespace(), "jobType", w.JobType(), "agentName", w.AgentName()) + h.logger.Infow("initial worker registered", "namespace", w.Namespace, "jobType", w.JobType, "agentName", w.AgentName) err := h.agentServer.PublishWorkerRegistered(context.Background(), agent.DefaultHandlerNamespace, &emptypb.Empty{}) if err != nil { - w.Logger().Errorw("failed to publish worker registered", err, "namespace", w.Namespace(), "jobType", w.JobType(), "agentName", w.AgentName()) + w.Logger().Errorw("failed to publish worker registered", err, "namespace", w.Namespace, "jobType", w.JobType, "agentName", w.AgentName) } } } -func (h *AgentHandler) HandleWorkerDeregister(w *agent.Worker) { +func (h *AgentHandler) deregisterWorker(w *agent.Worker) { h.mu.Lock() defer h.mu.Unlock() - key := workerKey{w.AgentName(), w.Namespace(), w.JobType()} + delete(h.workers, w.ID) + + key := workerKey{w.AgentName, w.Namespace, w.JobType} workers, ok := h.namespaceWorkers[key] if !ok { @@ -286,11 +299,11 @@ func (h *AgentHandler) HandleWorkerDeregister(w *agent.Worker) { if len(workers) > 1 { h.namespaceWorkers[key] = slices.Delete(workers, index, index+1) } else { - h.logger.Debugw("last worker deregistered", "namespace", w.Namespace(), "jobType", w.JobType(), "agentName", w.AgentName()) + h.logger.Debugw("last worker deregistered", "namespace", w.Namespace, "jobType", w.JobType, "agentName", w.AgentName) delete(h.namespaceWorkers, key) - topic := agent.GetAgentTopic(w.AgentName(), w.Namespace()) - if w.JobType() == livekit.JobType_JT_ROOM { + topic := agent.GetAgentTopic(w.AgentName, w.Namespace) + if w.JobType == livekit.JobType_JT_ROOM { h.roomKeyCount-- h.agentServer.DeregisterJobRequestTopic(topic, h.roomTopic) } else { @@ -299,10 +312,10 @@ func (h *AgentHandler) HandleWorkerDeregister(w *agent.Worker) { } // agentNames and namespaces contains repeated entries for each agentNames/namespaces combinations - if i := slices.Index(h.namespaces, w.Namespace()); i != -1 { + if i := slices.Index(h.namespaces, w.Namespace); i != -1 { h.namespaces = slices.Delete(h.namespaces, i, i+1) } - if i := slices.Index(h.agentNames, w.AgentName()); i != -1 { + if i := slices.Index(h.agentNames, w.AgentName); i != -1 { h.agentNames = slices.Delete(h.agentNames, i, i+1) } } @@ -344,7 +357,7 @@ func (h *AgentHandler) JobRequest(ctx context.Context, job *livekit.Job) (*rpc.J "jobID", job.Id, "namespace", job.Namespace, "agentName", job.AgentName, - "workerID", selected.ID(), + "workerID", selected.ID, } if job.Room != nil { values = append(values, "room", job.Room.Name, "roomID", job.Room.Sid) @@ -381,7 +394,7 @@ func (h *AgentHandler) JobRequestAffinity(ctx context.Context, job *livekit.Job) var affinity float32 for _, w := range h.workers { - if w.AgentName() != job.AgentName || w.Namespace() != job.Namespace || w.JobType() != job.Type { + if w.AgentName != job.AgentName || w.Namespace != job.Namespace || w.JobType != job.Type { continue } @@ -451,18 +464,11 @@ func (h *AgentHandler) selectWorkerWeightedByLoad(key workerKey, ignore map[*age return nil, errors.New("no workers available") } - normalizeLoad := func(load float32) int { - if load >= 1 { - return 0 - } - return int((1 - load) * 100) - } - - normalizedLoads := make(map[*agent.Worker]int) - var availableSum int + normalizedLoads := make(map[*agent.Worker]float32) + var availableSum float32 for _, w := range workers { if _, ok := ignore[w]; !ok && w.Status() == livekit.WorkerStatus_WS_AVAILABLE { - normalizedLoads[w] = normalizeLoad(w.Load()) + normalizedLoads[w] = max(0, 1-w.Load()) availableSum += normalizedLoads[w] } } @@ -471,14 +477,11 @@ func (h *AgentHandler) selectWorkerWeightedByLoad(key workerKey, ignore map[*age return nil, errors.New("no workers with sufficient capacity") } - threshold := rand.Intn(availableSum) - var currentSum int + currentSum := rand.Float32() * availableSum for w, load := range normalizedLoads { - currentSum += load - if currentSum >= threshold { + if currentSum -= load; currentSum <= 0 { return w, nil } } - - return nil, errors.New("no workers available") + return workers[0], nil } diff --git a/pkg/service/auth.go b/pkg/service/auth.go index 58e0ba7c1..188759fab 100644 --- a/pkg/service/auth.go +++ b/pkg/service/auth.go @@ -107,6 +107,13 @@ func (m *APIKeyAuthMiddleware) ServeHTTP(w http.ResponseWriter, r *http.Request, next.ServeHTTP(w, r) } +func WithAPIKey(ctx context.Context, grants *auth.ClaimGrants, apiKey string) context.Context { + return context.WithValue(ctx, grantsKey{}, &grantsValue{ + claims: grants, + apiKey: apiKey, + }) +} + func GetGrants(ctx context.Context) *auth.ClaimGrants { val := ctx.Value(grantsKey{}) v, ok := val.(*grantsValue) diff --git a/pkg/service/wsprotocol.go b/pkg/service/wsprotocol.go index cdfb72a8a..9f4b50e23 100644 --- a/pkg/service/wsprotocol.go +++ b/pkg/service/wsprotocol.go @@ -56,6 +56,10 @@ func (c *WSSignalConnection) Close() error { return c.conn.Close() } +func (c *WSSignalConnection) SetReadDeadline(deadline time.Time) error { + return c.conn.SetReadDeadline(deadline) +} + func (c *WSSignalConnection) ReadRequest() (*livekit.SignalRequest, int, error) { for { // handle special messages and pass on the rest