From 58ce5ed605dc8cb32168757f67f7104286a24396 Mon Sep 17 00:00:00 2001 From: Paul Wells Date: Wed, 18 Sep 2024 00:17:45 -0700 Subject: [PATCH 1/7] simplify agent registration --- pkg/agent/agent_test.go | 128 +++++------ pkg/agent/worker.go | 430 +++++++++++++++++++----------------- pkg/service/agentservice.go | 161 +++++++------- pkg/service/auth.go | 7 + pkg/service/wsprotocol.go | 4 + 5 files changed, 377 insertions(+), 353 deletions(-) diff --git a/pkg/agent/agent_test.go b/pkg/agent/agent_test.go index 52dee7839..670af6213 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,32 @@ 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))) + 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 +136,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/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/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 From fc87682e10109300d1687fc9a5789228835437c7 Mon Sep 17 00:00:00 2001 From: Paul Wells Date: Wed, 18 Sep 2024 00:34:40 -0700 Subject: [PATCH 2/7] testutils --- pkg/agent/testutils/server.go | 482 ++++++++++++++++++++++++++++++++++ 1 file changed, 482 insertions(+) create mode 100644 pkg/agent/testutils/server.go diff --git a/pkg/agent/testutils/server.go b/pkg/agent/testutils/server.go new file mode 100644 index 000000000..a8d48e3c5 --- /dev/null +++ b/pkg/agent/testutils/server.go @@ -0,0 +1,482 @@ +package testutils + +import ( + "context" + "errors" + "io" + "math" + "math/rand/v2" + "sync" + "time" + + "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/utils" + "github.com/livekit/protocol/utils/guid" + "github.com/livekit/protocol/utils/must" + "github.com/livekit/psrpc" +) + +type AgentService interface { + HandleConnection(context.Context, agent.SignalConn, agent.WorkerProtocolVersion) + DrainConnections(time.Duration) +} + +type TestServer struct { + AgentService +} + +func NewTestServer(bus psrpc.MessageBus) *TestServer { + return NewTestServerWithService(must.Get(service.NewAgentService( + &config.Config{Region: "test"}, + &livekit.Node{Id: guid.New("N_")}, + bus, + auth.NewSimpleKeyProvider("test", "verysecretsecret"), + ))) +} + +func NewTestServerWithService(s AgentService) *TestServer { + return &TestServer{s} +} + +type SimulatedWorkerOptions struct { + Context context.Context + Label string + SupportResume bool + DefaultJobLoad float32 + JobLoadThreshold float32 + DefaultWorkerLoad float32 + HandleAvailability func(AgentJobRequest) + HandleAssignment func(*livekit.Job) JobLoad +} + +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 + } +} + +func WithJobAssignmentHandler(h func(*livekit.Job) JobLoad) SimulatedWorkerOption { + return func(o *SimulatedWorkerOptions) { + o.HandleAssignment = h + } +} + +func WithJobLoad(l JobLoad) SimulatedWorkerOption { + return WithJobAssignmentHandler(func(j *livekit.Job) JobLoad { return l }) +} + +func WithDefaultWorkerLoad(load float32) SimulatedWorkerOption { + return func(o *SimulatedWorkerOptions) { + o.DefaultWorkerLoad = load + } +} + +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) + } + + w := &AgentWorker{ + workerMessages: make(chan *livekit.WorkerMessage, 1), + jobs: map[string]*AgentJob{}, + SimulatedWorkerOptions: o, + + RegisterWorkerResponses: utils.NewDefaultEventObserverList[*livekit.RegisterWorkerResponse](), + AvailabilityRequests: utils.NewDefaultEventObserverList[*livekit.AvailabilityRequest](), + JobAssignments: utils.NewDefaultEventObserverList[*livekit.JobAssignment](), + JobTerminations: utils.NewDefaultEventObserverList[*livekit.JobTermination](), + WorkerPongs: utils.NewDefaultEventObserverList[*livekit.WorkerPong](), + } + w.ctx, w.cancel = context.WithCancel(context.Background()) + + if o.DefaultWorkerLoad > 0.0 { + w.sendStatus() + } + + ctx := service.WithAPIKey(o.Context, &auth.ClaimGrants{}, "test") + go h.HandleConnection(ctx, w, agent.CurrentProtocol) + + return w +} + +func (h *TestServer) Close() { + h.DrainConnections(1) +} + +var _ agent.SignalConn = (*AgentWorker)(nil) + +type JobLoad interface { + Load() float32 +} + +type AgentJob struct { + *livekit.Job + JobLoad +} + +type AgentJobRequest struct { + w *AgentWorker + *livekit.AvailabilityRequest +} + +func (r AgentJobRequest) Accept() { + identity := guid.New("PI_") + r.w.SendAvailability(&livekit.AvailabilityResponse{ + JobId: r.Job.Id, + Available: true, + SupportsResume: r.w.SupportResume, + ParticipantName: identity, + ParticipantIdentity: identity, + }) +} + +func (r AgentJobRequest) Reject() { + r.w.SendAvailability(&livekit.AvailabilityResponse{ + JobId: r.Job.Id, + Available: false, + }) +} + +type AgentWorker struct { + *SimulatedWorkerOptions + + fuse core.Fuse + mu sync.Mutex + ctx context.Context + cancel context.CancelFunc + workerMessages chan *livekit.WorkerMessage + serverMessages deque.Deque[*livekit.ServerMessage] + jobs map[string]*AgentJob + + RegisterWorkerResponses *utils.EventObserverList[*livekit.RegisterWorkerResponse] + AvailabilityRequests *utils.EventObserverList[*livekit.AvailabilityRequest] + JobAssignments *utils.EventObserverList[*livekit.JobAssignment] + JobTerminations *utils.EventObserverList[*livekit.JobTermination] + WorkerPongs *utils.EventObserverList[*livekit.WorkerPong] +} + +func (w *AgentWorker) statusWorker() { + t := time.NewTicker(2 * time.Second) + defer t.Stop() + + for !w.fuse.IsBroken() { + w.sendStatus() + <-t.C + } +} + +func (w *AgentWorker) Close() error { + w.mu.Lock() + defer w.mu.Unlock() + w.fuse.Break() + return nil +} + +func (w *AgentWorker) SetReadDeadline(t time.Time) error { + w.mu.Lock() + defer w.mu.Unlock() + if !w.fuse.IsBroken() { + cancel := w.cancel + if t.IsZero() { + w.ctx, w.cancel = context.WithCancel(context.Background()) + } else { + w.ctx, w.cancel = context.WithDeadline(context.Background(), t) + } + cancel() + } + return nil +} + +func (w *AgentWorker) ReadWorkerMessage() (*livekit.WorkerMessage, int, error) { + for { + w.mu.Lock() + ctx := w.ctx + w.mu.Unlock() + + select { + case <-w.fuse.Watch(): + return nil, 0, io.EOF + case <-ctx.Done(): + if err := ctx.Err(); errors.Is(err, context.DeadlineExceeded) { + return nil, 0, err + } + case m := <-w.workerMessages: + return m, 0, nil + } + } +} + +func (w *AgentWorker) WriteServerMessage(m *livekit.ServerMessage) (int, error) { + w.mu.Lock() + defer w.mu.Unlock() + w.serverMessages.PushBack(m) + if w.serverMessages.Len() == 1 { + go w.handleServerMessages() + } + return 0, nil +} + +func (w *AgentWorker) handleServerMessages() { + w.mu.Lock() + for w.serverMessages.Len() != 0 { + m := w.serverMessages.Front() + w.mu.Unlock() + + switch m := m.Message.(type) { + case *livekit.ServerMessage_Register: + w.handleRegister(m.Register) + case *livekit.ServerMessage_Availability: + w.handleAvailability(m.Availability) + case *livekit.ServerMessage_Assignment: + w.handleAssignment(m.Assignment) + case *livekit.ServerMessage_Termination: + w.handleTermination(m.Termination) + case *livekit.ServerMessage_Pong: + w.handlePong(m.Pong) + } + + w.mu.Lock() + w.serverMessages.PopFront() + } + w.mu.Unlock() +} + +func (w *AgentWorker) handleRegister(m *livekit.RegisterWorkerResponse) { + w.RegisterWorkerResponses.Emit(m) +} + +func (w *AgentWorker) handleAvailability(m *livekit.AvailabilityRequest) { + w.AvailabilityRequests.Emit(m) + if w.HandleAvailability != nil { + w.HandleAvailability(AgentJobRequest{w, m}) + } else { + AgentJobRequest{w, m}.Accept() + } +} + +func (w *AgentWorker) handleAssignment(m *livekit.JobAssignment) { + w.JobAssignments.Emit(m) + + var load JobLoad + if w.HandleAssignment != nil { + load = w.HandleAssignment(m.Job) + } + + if load == nil { + load = NewStableJobLoad(w.DefaultJobLoad) + } + + w.mu.Lock() + defer w.mu.Unlock() + w.jobs[m.Job.Id] = &AgentJob{m.Job, load} +} + +func (w *AgentWorker) handleTermination(m *livekit.JobTermination) { + w.JobTerminations.Emit(m) + + w.mu.Lock() + defer w.mu.Unlock() + delete(w.jobs, m.JobId) +} + +func (w *AgentWorker) handlePong(m *livekit.WorkerPong) { + w.WorkerPongs.Emit(m) +} + +func (w *AgentWorker) sendMessage(m *livekit.WorkerMessage) { + select { + case <-w.fuse.Watch(): + case w.workerMessages <- m: + } +} + +func (w *AgentWorker) SendRegister(m *livekit.RegisterWorkerRequest) { + w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_Register{ + Register: m, + }}) +} + +func (w *AgentWorker) SendAvailability(m *livekit.AvailabilityResponse) { + w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_Availability{ + Availability: m, + }}) +} + +func (w *AgentWorker) SendUpdateWorker(m *livekit.UpdateWorkerStatus) { + w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_UpdateWorker{ + UpdateWorker: m, + }}) +} + +func (w *AgentWorker) SendUpdateJob(m *livekit.UpdateJobStatus) { + w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_UpdateJob{ + UpdateJob: m, + }}) +} + +func (w *AgentWorker) SendPing(m *livekit.WorkerPing) { + w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_Ping{ + Ping: m, + }}) +} + +func (w *AgentWorker) SendSimulateJob(m *livekit.SimulateJobRequest) { + w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_SimulateJob{ + SimulateJob: m, + }}) +} + +func (w *AgentWorker) SendMigrateJob(m *livekit.MigrateJobRequest) { + w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_MigrateJob{ + MigrateJob: m, + }}) +} + +func (w *AgentWorker) sendStatus() { + w.mu.Lock() + var load float32 + jobCount := len(w.jobs) + + if len(w.jobs) == 0 { + load = w.DefaultWorkerLoad + } else { + for _, j := range w.jobs { + load += j.Load() + } + } + w.mu.Unlock() + + status := livekit.WorkerStatus_WS_AVAILABLE + if load > w.JobLoadThreshold { + status = livekit.WorkerStatus_WS_FULL + } + + w.SendUpdateWorker(&livekit.UpdateWorkerStatus{ + Status: &status, + Load: load, + JobCount: int32(jobCount), + }) +} + +func (w *AgentWorker) Register(agentName string, jobType livekit.JobType) { + w.SendRegister(&livekit.RegisterWorkerRequest{ + Type: jobType, + AgentName: agentName, + }) + go w.statusWorker() +} + +func (w *AgentWorker) SimulateRoomJob(roomName string) { + w.SendSimulateJob(&livekit.SimulateJobRequest{ + Type: livekit.JobType_JT_ROOM, + Room: &livekit.Room{ + Sid: guid.New(guid.RoomPrefix), + Name: roomName, + }, + }) +} + +func (w *AgentWorker) Jobs() []*AgentJob { + w.mu.Lock() + defer w.mu.Unlock() + return maps.Values(w.jobs) +} + +type stableJobLoad struct { + load float32 +} + +func NewStableJobLoad(load float32) JobLoad { + return stableJobLoad{load} +} + +func (s stableJobLoad) Load() float32 { + return s.load +} + +type periodicJobLoad struct { + amplitude float64 + period time.Duration + epoch time.Time +} + +func NewPeriodicJobLoad(max float32, period time.Duration) JobLoad { + return periodicJobLoad{ + amplitude: float64(max / 2), + period: period, + epoch: time.Now().Add(-time.Duration(rand.Int64N(int64(period)))), + } +} + +func (s periodicJobLoad) Load() float32 { + a := math.Sin(time.Since(s.epoch).Seconds() / s.period.Seconds() * math.Pi * 2) + return float32(s.amplitude + a*s.amplitude) +} + +type uniformRandomJobLoad struct { + min, max float32 + rng func() float64 +} + +func NewUniformRandomJobLoad(min, max float32) JobLoad { + return uniformRandomJobLoad{min, max, rand.Float64} +} + +func NewUniformRandomJobLoadWithRNG(min, max float32, rng *rand.Rand) JobLoad { + return uniformRandomJobLoad{min, max, rng.Float64} +} + +func (s uniformRandomJobLoad) Load() float32 { + return rand.Float32()*(s.max-s.min) + s.min +} + +type normalRandomJobLoad struct { + mean, stddev float64 + rng func() float64 +} + +func NewNormalRandomJobLoad(mean, stddev float64) JobLoad { + return normalRandomJobLoad{mean, stddev, rand.Float64} +} + +func NewNormalRandomJobLoadWithRNG(mean, stddev float64, rng *rand.Rand) JobLoad { + return normalRandomJobLoad{mean, stddev, rng.Float64} +} + +func (s normalRandomJobLoad) Load() float32 { + u := 1 - s.rng() + v := s.rng() + z := math.Sqrt(-2*math.Log(u)) * math.Cos(2*math.Pi*v) + return float32(max(0, z*s.stddev+s.mean)) +} From 81f5f2a225278b83d9d875ff9f956fc4cb39a3bc Mon Sep 17 00:00:00 2001 From: Paul Wells Date: Wed, 18 Sep 2024 00:35:19 -0700 Subject: [PATCH 3/7] deps --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) 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= From aa27600638211b32d10971253bc2f35f7761adc0 Mon Sep 17 00:00:00 2001 From: Paul Wells Date: Wed, 18 Sep 2024 00:49:44 -0700 Subject: [PATCH 4/7] fix --- pkg/rtc/types/interfaces.go | 1 + 1 file changed, 1 insertion(+) diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index 31f7196c3..b2059256d 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 } From 6ad48b0cc9b39118e8a8e7f977cafb5266196feb Mon Sep 17 00:00:00 2001 From: Paul Wells Date: Wed, 18 Sep 2024 00:58:48 -0700 Subject: [PATCH 5/7] gen --- .../types/typesfakes/fake_websocket_client.go | 74 +++++++++++++++++++ 1 file changed, 74 insertions(+) diff --git a/pkg/rtc/types/typesfakes/fake_websocket_client.go b/pkg/rtc/types/typesfakes/fake_websocket_client.go index 0a9c14b7f..93d0957ad 100644 --- a/pkg/rtc/types/typesfakes/fake_websocket_client.go +++ b/pkg/rtc/types/typesfakes/fake_websocket_client.go @@ -33,6 +33,17 @@ type FakeWebsocketClient struct { result2 []byte result3 error } + SetReadDeadlineStub func(time.Time) error + setReadDeadlineMutex sync.RWMutex + setReadDeadlineArgsForCall []struct { + arg1 time.Time + } + setReadDeadlineReturns struct { + result1 error + } + setReadDeadlineReturnsOnCall map[int]struct { + result1 error + } WriteControlStub func(int, []byte, time.Time) error writeControlMutex sync.RWMutex writeControlArgsForCall []struct { @@ -174,6 +185,67 @@ func (fake *FakeWebsocketClient) ReadMessageReturnsOnCall(i int, result1 int, re }{result1, result2, result3} } +func (fake *FakeWebsocketClient) SetReadDeadline(arg1 time.Time) error { + fake.setReadDeadlineMutex.Lock() + ret, specificReturn := fake.setReadDeadlineReturnsOnCall[len(fake.setReadDeadlineArgsForCall)] + fake.setReadDeadlineArgsForCall = append(fake.setReadDeadlineArgsForCall, struct { + arg1 time.Time + }{arg1}) + stub := fake.SetReadDeadlineStub + fakeReturns := fake.setReadDeadlineReturns + fake.recordInvocation("SetReadDeadline", []interface{}{arg1}) + fake.setReadDeadlineMutex.Unlock() + if stub != nil { + return stub(arg1) + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeWebsocketClient) SetReadDeadlineCallCount() int { + fake.setReadDeadlineMutex.RLock() + defer fake.setReadDeadlineMutex.RUnlock() + return len(fake.setReadDeadlineArgsForCall) +} + +func (fake *FakeWebsocketClient) SetReadDeadlineCalls(stub func(time.Time) error) { + fake.setReadDeadlineMutex.Lock() + defer fake.setReadDeadlineMutex.Unlock() + fake.SetReadDeadlineStub = stub +} + +func (fake *FakeWebsocketClient) SetReadDeadlineArgsForCall(i int) time.Time { + fake.setReadDeadlineMutex.RLock() + defer fake.setReadDeadlineMutex.RUnlock() + argsForCall := fake.setReadDeadlineArgsForCall[i] + return argsForCall.arg1 +} + +func (fake *FakeWebsocketClient) SetReadDeadlineReturns(result1 error) { + fake.setReadDeadlineMutex.Lock() + defer fake.setReadDeadlineMutex.Unlock() + fake.SetReadDeadlineStub = nil + fake.setReadDeadlineReturns = struct { + result1 error + }{result1} +} + +func (fake *FakeWebsocketClient) SetReadDeadlineReturnsOnCall(i int, result1 error) { + fake.setReadDeadlineMutex.Lock() + defer fake.setReadDeadlineMutex.Unlock() + fake.SetReadDeadlineStub = nil + if fake.setReadDeadlineReturnsOnCall == nil { + fake.setReadDeadlineReturnsOnCall = make(map[int]struct { + result1 error + }) + } + fake.setReadDeadlineReturnsOnCall[i] = struct { + result1 error + }{result1} +} + func (fake *FakeWebsocketClient) WriteControl(arg1 int, arg2 []byte, arg3 time.Time) error { var arg2Copy []byte if arg2 != nil { @@ -316,6 +388,8 @@ func (fake *FakeWebsocketClient) Invocations() map[string][][]interface{} { defer fake.closeMutex.RUnlock() fake.readMessageMutex.RLock() defer fake.readMessageMutex.RUnlock() + fake.setReadDeadlineMutex.RLock() + defer fake.setReadDeadlineMutex.RUnlock() fake.writeControlMutex.RLock() defer fake.writeControlMutex.RUnlock() fake.writeMessageMutex.RLock() From d09a3386093c8edd8f087a991ec110d6dbb00c84 Mon Sep 17 00:00:00 2001 From: Paul Wells Date: Wed, 18 Sep 2024 01:12:18 -0700 Subject: [PATCH 6/7] cleanup --- pkg/agent/testutil/server.go | 495 ----------------------------------- 1 file changed, 495 deletions(-) delete mode 100644 pkg/agent/testutil/server.go diff --git a/pkg/agent/testutil/server.go b/pkg/agent/testutil/server.go deleted file mode 100644 index 04f437534..000000000 --- a/pkg/agent/testutil/server.go +++ /dev/null @@ -1,495 +0,0 @@ -package testutil - -import ( - "context" - "errors" - "io" - "math" - "math/rand/v2" - "sync" - "time" - - "github.com/frostbyte73/core" - "github.com/gammazero/deque" - - "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/psrpc" -) - -type TestServer struct { - *service.AgentService - keyProvider auth.KeyProvider -} - -func NewTestServer(bus psrpc.MessageBus) *TestServer { - keyProvider := auth.NewSimpleKeyProvider("test", "verysecretsecret") - - s := must.Get(service.NewAgentService( - &config.Config{Region: "test"}, - &livekit.Node{Id: guid.New("N_")}, - bus, - keyProvider, - )) - - return &TestServer{ - AgentService: s, - keyProvider: keyProvider, - } -} - -type SimulatedWorkerOptions struct { - SupportResume bool - DefaultJobLoad float32 - JobLoadThreshold float32 - DefaultWorkerLoad float32 - HandleAvailability func(AgentJobRequest) - HandleAssignment func(*livekit.Job) JobLoad -} - -type SimulatedWorkerOption func(*SimulatedWorkerOptions) - -func WithJobAvailabilityHandler(h func(AgentJobRequest)) SimulatedWorkerOption { - return func(o *SimulatedWorkerOptions) { - o.HandleAvailability = h - } -} - -func WithJobAssignmentHandler(h func(*livekit.Job) JobLoad) SimulatedWorkerOption { - return func(o *SimulatedWorkerOptions) { - o.HandleAssignment = h - } -} - -func WithJobLoad(l JobLoad) SimulatedWorkerOption { - return WithJobAssignmentHandler(func(j *livekit.Job) JobLoad { return l }) -} - -func WithDefaultWorkerLoad(load float32) SimulatedWorkerOption { - return func(o *SimulatedWorkerOptions) { - o.DefaultWorkerLoad = load - } -} - -func (h *TestServer) SimulateAgentWorker(opts ...SimulatedWorkerOption) *AgentWorker { - o := &SimulatedWorkerOptions{ - 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) - } - - w := &AgentWorker{ - workerMessages: make(chan *livekit.WorkerMessage, 1), - jobs: map[string]*AgentJob{}, - SimulatedWorkerOptions: o, - - RegisterWorkerResponses: utils.NewDefaultEventObserverList[*livekit.RegisterWorkerResponse](), - AvailabilityRequests: utils.NewDefaultEventObserverList[*livekit.AvailabilityRequest](), - JobAssignments: utils.NewDefaultEventObserverList[*livekit.JobAssignment](), - JobTerminations: utils.NewDefaultEventObserverList[*livekit.JobTermination](), - WorkerPongs: utils.NewDefaultEventObserverList[*livekit.WorkerPong](), - } - w.ctx, w.cancel = context.WithCancel(context.Background()) - - if o.DefaultWorkerLoad > 0.0 { - w.sendStatus() - } - - go w.worker() - go h.handleConnection(w) - 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() - } -} - -var _ agent.SignalConn = (*AgentWorker)(nil) - -type JobLoad interface { - Load() float32 -} - -type AgentJob struct { - *livekit.Job - JobLoad -} - -type AgentJobRequest struct { - w *AgentWorker - *livekit.AvailabilityRequest -} - -func (r AgentJobRequest) Accept() { - identity := guid.New("PI_") - r.w.SendAvailability(&livekit.AvailabilityResponse{ - JobId: r.Job.Id, - Available: true, - SupportsResume: r.w.SupportResume, - ParticipantName: identity, - ParticipantIdentity: identity, - }) -} - -func (r AgentJobRequest) Reject() { - r.w.SendAvailability(&livekit.AvailabilityResponse{ - JobId: r.Job.Id, - Available: false, - }) -} - -type AgentWorker struct { - Name string - *SimulatedWorkerOptions - - fuse core.Fuse - mu sync.Mutex - ctx context.Context - cancel context.CancelFunc - workerMessages chan *livekit.WorkerMessage - serverMessages deque.Deque[*livekit.ServerMessage] - jobs map[string]*AgentJob - - RegisterWorkerResponses *utils.EventObserverList[*livekit.RegisterWorkerResponse] - AvailabilityRequests *utils.EventObserverList[*livekit.AvailabilityRequest] - JobAssignments *utils.EventObserverList[*livekit.JobAssignment] - JobTerminations *utils.EventObserverList[*livekit.JobTermination] - WorkerPongs *utils.EventObserverList[*livekit.WorkerPong] -} - -func (w *AgentWorker) worker() { - t := time.NewTicker(5 * time.Second) - defer t.Stop() - - for !w.fuse.IsBroken() { - <-t.C - w.sendStatus() - } -} - -func (w *AgentWorker) Close() error { - w.mu.Lock() - defer w.mu.Unlock() - w.fuse.Break() - return nil -} - -func (w *AgentWorker) SetReadDeadline(t time.Time) error { - w.mu.Lock() - defer w.mu.Unlock() - if !w.fuse.IsBroken() { - cancel := w.cancel - if t.IsZero() { - w.ctx, w.cancel = context.WithCancel(context.Background()) - } else { - w.ctx, w.cancel = context.WithDeadline(context.Background(), t) - } - cancel() - } - return nil -} - -func (w *AgentWorker) ReadWorkerMessage() (*livekit.WorkerMessage, int, error) { - for { - w.mu.Lock() - ctx := w.ctx - w.mu.Unlock() - - select { - case <-w.fuse.Watch(): - return nil, 0, io.EOF - case <-ctx.Done(): - if err := ctx.Err(); errors.Is(err, context.DeadlineExceeded) { - return nil, 0, err - } - case m := <-w.workerMessages: - return m, 0, nil - } - } -} - -func (w *AgentWorker) WriteServerMessage(m *livekit.ServerMessage) (int, error) { - w.mu.Lock() - defer w.mu.Unlock() - w.serverMessages.PushBack(m) - if w.serverMessages.Len() == 1 { - go w.handleServerMessages() - } - return 0, nil -} - -func (w *AgentWorker) handleServerMessages() { - w.mu.Lock() - for w.serverMessages.Len() != 0 { - m := w.serverMessages.Front() - w.mu.Unlock() - - switch m := m.Message.(type) { - case *livekit.ServerMessage_Register: - w.handleRegister(m.Register) - case *livekit.ServerMessage_Availability: - w.handleAvailability(m.Availability) - case *livekit.ServerMessage_Assignment: - w.handleAssignment(m.Assignment) - case *livekit.ServerMessage_Termination: - w.handleTermination(m.Termination) - case *livekit.ServerMessage_Pong: - w.handlePong(m.Pong) - } - - w.mu.Lock() - w.serverMessages.PopFront() - } - w.mu.Unlock() -} - -func (w *AgentWorker) handleRegister(m *livekit.RegisterWorkerResponse) { - w.RegisterWorkerResponses.Emit(m) -} - -func (w *AgentWorker) handleAvailability(m *livekit.AvailabilityRequest) { - w.AvailabilityRequests.Emit(m) - if w.HandleAvailability != nil { - w.HandleAvailability(AgentJobRequest{w, m}) - } else { - AgentJobRequest{w, m}.Accept() - } -} - -func (w *AgentWorker) handleAssignment(m *livekit.JobAssignment) { - m.Job.AgentName = w.Name - w.JobAssignments.Emit(m) - - var load JobLoad - if w.HandleAssignment != nil { - load = w.HandleAssignment(m.Job) - } - - if load == nil { - load = NewStableJobLoad(w.DefaultJobLoad) - } - - w.mu.Lock() - defer w.mu.Unlock() - w.jobs[m.Job.Id] = &AgentJob{m.Job, load} -} - -func (w *AgentWorker) handleTermination(m *livekit.JobTermination) { - w.JobTerminations.Emit(m) - - w.mu.Lock() - defer w.mu.Unlock() - delete(w.jobs, m.JobId) -} - -func (w *AgentWorker) handlePong(m *livekit.WorkerPong) { - w.WorkerPongs.Emit(m) -} - -func (w *AgentWorker) sendMessage(m *livekit.WorkerMessage) { - select { - case <-w.fuse.Watch(): - case w.workerMessages <- m: - } -} - -func (w *AgentWorker) SendRegister(m *livekit.RegisterWorkerRequest) { - w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_Register{ - Register: m, - }}) -} - -func (w *AgentWorker) SendAvailability(m *livekit.AvailabilityResponse) { - w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_Availability{ - Availability: m, - }}) -} - -func (w *AgentWorker) SendUpdateWorker(m *livekit.UpdateWorkerStatus) { - w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_UpdateWorker{ - UpdateWorker: m, - }}) -} - -func (w *AgentWorker) SendUpdateJob(m *livekit.UpdateJobStatus) { - w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_UpdateJob{ - UpdateJob: m, - }}) -} - -func (w *AgentWorker) SendPing(m *livekit.WorkerPing) { - w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_Ping{ - Ping: m, - }}) -} - -func (w *AgentWorker) SendSimulateJob(m *livekit.SimulateJobRequest) { - w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_SimulateJob{ - SimulateJob: m, - }}) -} - -func (w *AgentWorker) SendMigrateJob(m *livekit.MigrateJobRequest) { - w.sendMessage(&livekit.WorkerMessage{Message: &livekit.WorkerMessage_MigrateJob{ - MigrateJob: m, - }}) -} - -func (w *AgentWorker) sendStatus() { - w.mu.Lock() - var load float32 - jobCount := len(w.jobs) - - if len(w.jobs) == 0 { - load = w.DefaultWorkerLoad - } else { - for _, j := range w.jobs { - load += j.Load() - } - } - w.mu.Unlock() - - status := livekit.WorkerStatus_WS_AVAILABLE - if load > w.JobLoadThreshold { - status = livekit.WorkerStatus_WS_FULL - } - - w.SendUpdateWorker(&livekit.UpdateWorkerStatus{ - Status: &status, - Load: load, - JobCount: int32(jobCount), - }) -} - -func (w *AgentWorker) Register(name string, namespace string, jobType livekit.JobType) { - w.Name = name - w.SendRegister(&livekit.RegisterWorkerRequest{ - Type: jobType, - Namespace: &namespace, - }) - w.sendStatus() -} - -func (w *AgentWorker) SimulateRoomJob(roomName string) { - w.SendSimulateJob(&livekit.SimulateJobRequest{ - Type: livekit.JobType_JT_ROOM, - Room: &livekit.Room{ - Sid: guid.New(guid.RoomPrefix), - Name: roomName, - }, - }) -} - -type stableJobLoad struct { - load float32 -} - -func NewStableJobLoad(load float32) JobLoad { - return stableJobLoad{load} -} - -func (s stableJobLoad) Load() float32 { - return s.load -} - -type periodicJobLoad struct { - amplitude float64 - period time.Duration - epoch time.Time -} - -func NewPeriodicJobLoad(max float32, period time.Duration) JobLoad { - return periodicJobLoad{ - amplitude: float64(max / 2), - period: period, - epoch: time.Now().Add(-time.Duration(rand.Int64N(int64(period)))), - } -} - -func (s periodicJobLoad) Load() float32 { - a := math.Sin(time.Since(s.epoch).Seconds() / s.period.Seconds() * math.Pi * 2) - return float32(s.amplitude + a*s.amplitude) -} - -type uniformRandomJobLoad struct { - min, max float32 - rng func() float64 -} - -func NewUniformRandomJobLoad(min, max float32) JobLoad { - return uniformRandomJobLoad{min, max, rand.Float64} -} - -func NewUniformRandomJobLoadWithRNG(min, max float32, rng *rand.Rand) JobLoad { - return uniformRandomJobLoad{min, max, rng.Float64} -} - -func (s uniformRandomJobLoad) Load() float32 { - return rand.Float32()*(s.max-s.min) + s.min -} - -type normalRandomJobLoad struct { - mean, stddev float64 - rng func() float64 -} - -func NewNormalRandomJobLoad(mean, stddev float64) JobLoad { - return normalRandomJobLoad{mean, stddev, rand.Float64} -} - -func NewNormalRandomJobLoadWithRNG(mean, stddev float64, rng *rand.Rand) JobLoad { - return normalRandomJobLoad{mean, stddev, rng.Float64} -} - -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) - return float32(max(0, z*s.stddev+s.mean)) -} From eed840b5359590d689a6f66834ad3bf8a06f22e0 Mon Sep 17 00:00:00 2001 From: Paul Wells Date: Wed, 18 Sep 2024 01:28:02 -0700 Subject: [PATCH 7/7] lower job load --- pkg/agent/agent_test.go | 5 ++++- pkg/agent/testutils/server.go | 5 ++--- 2 files changed, 6 insertions(+), 4 deletions(-) diff --git a/pkg/agent/agent_test.go b/pkg/agent/agent_test.go index 670af6213..079cf8e02 100644 --- a/pkg/agent/agent_test.go +++ b/pkg/agent/agent_test.go @@ -106,7 +106,10 @@ func TestAgentLoadBalancing(t *testing.T) { agents := make([]*testutils.AgentWorker, totalWorkers) for i := 0; i < totalWorkers; i++ { - agents[i] = server.SimulateAgentWorker(testutils.WithLabel(fmt.Sprintf("agent-%d", i))) + 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) } diff --git a/pkg/agent/testutils/server.go b/pkg/agent/testutils/server.go index a8d48e3c5..4178e0100 100644 --- a/pkg/agent/testutils/server.go +++ b/pkg/agent/testutils/server.go @@ -21,6 +21,7 @@ import ( "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" ) @@ -103,9 +104,7 @@ func (h *TestServer) SimulateAgentWorker(opts ...SimulatedWorkerOption) *AgentWo 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),