Merge remote-tracking branch 'origin/agents-cleanup' into agents-temp

This commit is contained in:
Paul Wells
2024-09-18 01:58:03 -07:00
9 changed files with 433 additions and 419 deletions
+1 -1
View File
@@ -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
+2 -2
View File
@@ -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=
+61 -70
View File
@@ -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)
})
}
@@ -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))
}
+226 -204
View File
@@ -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
}
+1
View File
@@ -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
}
+82 -79
View File
@@ -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
}
+7
View File
@@ -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)
+4
View File
@@ -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