Files
livekit/pkg/agent/agent_test.go
T
Paul Wells 1ba2cadccd agent: release JobTerminate handler when job ends before registration (#4908)
* agent: release JobTerminate handler when job ends before registration

JobRequest registered the job's JobTerminate handler after AssignJob
returned, so an ended UpdateJobStatus or worker disconnect handled in
between deregistered nothing and the handler leaked. AssignJob also sent
the assignment before recording the job as running, so an immediate
ended update was dropped.

Record the job before sending the assignment, and after registering the
handler recheck under h.mu that the job is still running and its worker
still registered, deregistering otherwise.

Fixes #4901

* replace utils.CloneProto with proto.CloneOf

* agent: wait on worker readiness and poll for released handlers in test
2026-09-25 00:58:22 -07:00

369 lines
11 KiB
Go

package agent_test
import (
"context"
"errors"
"fmt"
"sync"
"testing"
"time"
"github.com/stretchr/testify/require"
"go.uber.org/atomic"
"github.com/livekit/livekit-server/pkg/agent"
"github.com/livekit/livekit-server/pkg/agent/testutils"
"github.com/livekit/protocol/auth"
"github.com/livekit/protocol/livekit"
"github.com/livekit/protocol/rpc"
"github.com/livekit/protocol/utils/guid"
"github.com/livekit/protocol/utils/must"
"github.com/livekit/psrpc"
)
func TestAgent(t *testing.T) {
testAgentName := "test_agent"
t.Run("dispatched jobs are assigned to a worker", func(t *testing.T) {
bus := psrpc.NewLocalMessageBus()
client := must.Get(rpc.NewAgentInternalClient(bus))
server := testutils.NewTestServer(bus)
t.Cleanup(server.Close)
worker := server.SimulateAgentWorker()
worker.Register(testAgentName, livekit.JobType_JT_ROOM)
jobAssignments := worker.JobAssignments.Observe()
job := &livekit.Job{
Id: guid.New(guid.AgentJobPrefix),
DispatchId: guid.New(guid.AgentDispatchPrefix),
Type: livekit.JobType_JT_ROOM,
Room: &livekit.Room{},
AgentName: testAgentName,
}
_, err := client.JobRequest(context.Background(), testAgentName, agent.RoomAgentTopic, job)
require.NoError(t, err)
select {
case a := <-jobAssignments.Events():
require.EqualValues(t, job.Id, a.Job.Id)
v, err := auth.ParseAPIToken(a.Token)
require.NoError(t, err)
_, claims, err := v.Verify(server.TestAPISecret)
require.NoError(t, err)
require.Equal(t, testAgentName, claims.Attributes[agent.AgentNameAttributeKey])
case <-time.After(time.Second):
require.Fail(t, "job assignment timeout")
}
})
}
func testBatchJobRequest(t require.TestingT, 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)
}
}
}
}()
}
// wait for agent registration
time.Sleep(100 * time.Millisecond)
var wg sync.WaitGroup
for i := 0; i < totalJobs; i += batchSize {
wg.Add(1)
go func(start int) {
defer wg.Done()
for j := start; j < start+batchSize && j < totalJobs; j++ {
job := &livekit.Job{
Id: guid.New(guid.AgentJobPrefix),
DispatchId: guid.New(guid.AgentDispatchPrefix),
Type: livekit.JobType_JT_ROOM,
Room: &livekit.Room{},
AgentName: "test",
}
_, err := client.JobRequest(context.Background(), "test", agent.RoomAgentTopic, job)
require.NoError(t, err)
}
}(i)
}
wg.Wait()
return done
}
func TestAgentLoadBalancing(t *testing.T) {
t.Run("jobs are distributed normally with baseline worker load", func(t *testing.T) {
totalWorkers := 5
totalJobs := 100
bus := psrpc.NewLocalMessageBus()
client := must.Get(rpc.NewAgentInternalClient(bus))
t.Cleanup(client.Close)
server := testutils.NewTestServer(bus)
t.Cleanup(server.Close)
agents := make([]*testutils.AgentWorker, totalWorkers)
for i := range totalWorkers {
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)
}
select {
case <-testBatchJobRequest(t, 10, totalJobs, client, agents):
case <-time.After(time.Second):
require.Fail(t, "job assignment timeout")
}
jobCount := make(map[string]int)
for _, w := range agents {
jobCount[w.Label] = len(w.Jobs())
}
// check that jobs are distributed normally
for i := range totalWorkers {
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
}
})
t.Run("jobs are distributed with variable and overloaded worker load", func(t *testing.T) {
totalWorkers := 4
totalJobs := 15
bus := psrpc.NewLocalMessageBus()
client := must.Get(rpc.NewAgentInternalClient(bus))
t.Cleanup(client.Close)
server := testutils.NewTestServer(bus)
t.Cleanup(server.Close)
agents := make([]*testutils.AgentWorker, totalWorkers)
for i := range totalWorkers {
label := fmt.Sprintf("agent-%d", i)
if i%2 == 0 {
// make sure we have some workers that can accept jobs
agents[i] = server.SimulateAgentWorker(testutils.WithLabel(label))
} else {
agents[i] = server.SimulateAgentWorker(testutils.WithLabel(label), testutils.WithDefaultWorkerLoad(0.9))
}
agents[i].Register("test", livekit.JobType_JT_ROOM)
}
select {
case <-testBatchJobRequest(t, 1, totalJobs, client, agents):
case <-time.After(time.Second):
require.Fail(t, "job assignment timeout")
}
jobCount := make(map[string]int)
for _, w := range agents {
jobCount[w.Label] = len(w.Jobs())
}
for i := range totalWorkers {
label := fmt.Sprintf("agent-%d", i)
if i%2 == 0 {
require.GreaterOrEqual(t, jobCount[label], 2)
} else {
require.Equal(t, 0, jobCount[label])
}
require.GreaterOrEqual(t, jobCount[label], 0)
}
})
}
func TestConnectionClosedOnDispatchError(t *testing.T) {
t.Run("connection closed when unknown message type received", func(t *testing.T) {
bus := psrpc.NewLocalMessageBus()
server := testutils.NewTestServer(bus)
t.Cleanup(server.Close)
// register agent
worker := server.SimulateAgentWorker()
worker.Register("test_agent", livekit.JobType_JT_ROOM)
responses := worker.RegisterWorkerResponses.Observe()
select {
case <-responses.Events():
// registered
case <-time.After(time.Second):
require.Fail(t, "registration timeout")
}
responses.Stop()
// send invalid message (nil Message field triggers ErrUnknownWorkerSignal)
worker.SendMessage(&livekit.WorkerMessage{Message: nil})
select {
case <-worker.Closed():
// connection closed
case <-time.After(time.Second):
require.Fail(t, "connection should have been closed after dispatch error")
}
})
}
func TestDrainConnectionsDoesNotDeadlock(t *testing.T) {
for _, force := range []bool{false, true} {
t.Run(fmt.Sprintf("force=%v", force), func(t *testing.T) {
bus := psrpc.NewLocalMessageBus()
server := testutils.NewTestServer(bus)
t.Cleanup(server.Close)
worker := server.SimulateAgentWorker()
worker.Register("drain_agent", livekit.JobType_JT_ROOM)
responses := worker.RegisterWorkerResponses.Observe()
select {
case <-responses.Events():
case <-time.After(time.Second):
require.Fail(t, "registration timeout")
}
responses.Stop()
done := make(chan struct{})
go func() {
server.DrainConnections(time.Millisecond, force)
close(done)
}()
select {
case <-done:
case <-time.After(5 * time.Second):
require.Fail(t, "DrainConnections deadlocked while closing workers")
}
select {
case <-worker.Closed():
case <-time.After(time.Second):
require.Fail(t, "worker should be closed after drain")
}
})
}
}
func TestJobTerminateHandlerReleased(t *testing.T) {
const jobCount = 50
requestJobs := func(t *testing.T, client rpc.AgentInternalClient, agentName func(i int) string) []string {
jobIDs := make([]string, jobCount)
errs := make([]error, jobCount)
var wg sync.WaitGroup
for i := range jobCount {
job := &livekit.Job{
Id: guid.New(guid.AgentJobPrefix),
DispatchId: guid.New(guid.AgentDispatchPrefix),
Type: livekit.JobType_JT_ROOM,
Room: &livekit.Room{},
AgentName: agentName(i),
}
jobIDs[i] = job.Id
wg.Add(1)
go func() {
defer wg.Done()
_, errs[i] = client.JobRequest(context.Background(), job.AgentName, agent.RoomAgentTopic, job)
}()
}
wg.Wait()
for _, err := range errs {
require.NoError(t, err)
}
return jobIDs
}
// the server reads the ping only after registering the worker's job request topic.
waitRegistered := func(t *testing.T, w *testutils.AgentWorker, agentName string) {
pongs := w.WorkerPongs.Observe()
defer pongs.Stop()
w.Register(agentName, livekit.JobType_JT_ROOM)
w.SendPing(&livekit.WorkerPing{Timestamp: time.Now().UnixMilli()})
select {
case <-pongs.Events():
case <-time.After(5 * time.Second):
require.Fail(t, "registration timeout")
}
}
// any response means a JobTerminate handler is still registered.
countAnswered := func(client rpc.AgentInternalClient, jobIDs []string) int32 {
var answered atomic.Int32
var wg sync.WaitGroup
for _, id := range jobIDs {
wg.Add(1)
go func() {
defer wg.Done()
_, err := client.JobTerminate(context.Background(), id, &rpc.JobTerminateRequest{JobId: id}, psrpc.WithRequestTimeout(200*time.Millisecond))
if !errors.Is(err, psrpc.ErrRequestTimedOut) && !errors.Is(err, psrpc.ErrNoResponse) {
answered.Inc()
}
}()
}
wg.Wait()
return answered.Load()
}
// handlers are released as the server processes the worker's messages; a leaked one never is.
requireHandlersReleased := func(t *testing.T, client rpc.AgentInternalClient, jobIDs []string) {
var answered int32
require.Eventually(t, func() bool {
answered = countAnswered(client, jobIDs)
return answered == 0
}, 10*time.Second, 100*time.Millisecond, "jobs with leaked JobTerminate handlers: %d", answered)
}
t.Run("job fails on assignment", func(t *testing.T) {
bus := psrpc.NewLocalMessageBus()
client := must.Get(rpc.NewAgentInternalClient(bus))
server := testutils.NewTestServer(bus)
t.Cleanup(server.Close)
var worker *testutils.AgentWorker
worker = server.SimulateAgentWorker(testutils.WithJobAssignmentHandler(func(j *livekit.Job) testutils.JobLoad {
worker.SendUpdateJob(&livekit.UpdateJobStatus{JobId: j.Id, Status: livekit.JobStatus_JS_FAILED})
return testutils.NewStableJobLoad(0)
}))
waitRegistered(t, worker, "fail_agent")
jobIDs := requestJobs(t, client, func(int) string { return "fail_agent" })
requireHandlersReleased(t, client, jobIDs)
})
t.Run("worker disconnects on assignment", func(t *testing.T) {
bus := psrpc.NewLocalMessageBus()
client := must.Get(rpc.NewAgentInternalClient(bus))
server := testutils.NewTestServer(bus)
t.Cleanup(server.Close)
agentName := func(i int) string { return fmt.Sprintf("disconnect_agent_%d", i) }
for i := range jobCount {
var worker *testutils.AgentWorker
worker = server.SimulateAgentWorker(testutils.WithJobAssignmentHandler(func(j *livekit.Job) testutils.JobLoad {
worker.Close()
return nil
}))
waitRegistered(t, worker, agentName(i))
}
jobIDs := requestJobs(t, client, agentName)
requireHandlersReleased(t, client, jobIDs)
})
}