diff --git a/pkg/config/config.go b/pkg/config/config.go index cd39829b1..d00ae8060 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -10,8 +10,6 @@ import ( "github.com/pkg/errors" "github.com/urfave/cli/v2" "gopkg.in/yaml.v3" - - "github.com/livekit/livekit-server/pkg/routing/selector" ) var DefaultStunServers = []string{ @@ -112,9 +110,17 @@ type WebHookConfig struct { } type NodeSelectorConfig struct { - Kind string `yaml:"kind"` - SysloadLimit float32 `yaml:"sysload_limit"` - Regions []selector.RegionConfig `yaml:"regions"` + Kind string `yaml:"kind"` + SysloadLimit float32 `yaml:"sysload_limit"` + Regions []RegionConfig `yaml:"regions"` +} + +// RegionConfig lists available regions and their latitude/longitude, so the selector would prefer +// regions that are closer +type RegionConfig struct { + Name string `yaml:"name"` + Lat float64 `yaml:"lat"` + Lon float64 `yaml:"lon"` } func NewConfig(confString string, c *cli.Context) (*Config, error) { @@ -255,7 +261,3 @@ func (conf *Config) unmarshalKeys(keys string) error { } return nil } - -func GetAudioConfig(conf *Config) AudioConfig { - return conf.Audio -} diff --git a/pkg/routing/interfaces.go b/pkg/routing/interfaces.go index b8058a28d..2fe71929a 100644 --- a/pkg/routing/interfaces.go +++ b/pkg/routing/interfaces.go @@ -3,8 +3,12 @@ package routing import ( "context" + "github.com/go-redis/redis/v8" + "github.com/livekit/protocol/logger" livekit "github.com/livekit/protocol/proto" "google.golang.org/protobuf/proto" + + "github.com/livekit/livekit-server/pkg/config" ) //go:generate go run github.com/maxbrunsfeld/counterfeiter/v6 -generate @@ -40,15 +44,17 @@ type RTCMessageCallback func(ctx context.Context, roomName, identity string, msg // Router allows multiple nodes to coordinate the participant session //counterfeiter:generate . Router type Router interface { - GetNodeForRoom(ctx context.Context, roomName string) (*livekit.Node, error) - SetNodeForRoom(ctx context.Context, roomName string, nodeId string) error - ClearRoomState(ctx context.Context, roomName string) error RegisterNode() error UnregisterNode() error RemoveDeadNodes() error + GetNode(nodeId string) (*livekit.Node, error) ListNodes() ([]*livekit.Node, error) + GetNodeForRoom(ctx context.Context, roomName string) (*livekit.Node, error) + SetNodeForRoom(ctx context.Context, roomName, nodeId string) error + ClearRoomState(ctx context.Context, roomName string) error + // StartParticipantSignal participant signal connection is ready to start StartParticipantSignal(ctx context.Context, roomName string, pi ParticipantInit) (connectionId string, reqSink MessageSink, resSource MessageSource, err error) @@ -62,12 +68,16 @@ type Router interface { OnRTCMessage(callback RTCMessageCallback) Start() error - PreStop() + Drain() Stop() } -// NodeSelector selects an appropriate node to run the current session -//counterfeiter:generate . NodeSelector -type NodeSelector interface { - SelectNode(nodes []*livekit.Node, room *livekit.Room) (*livekit.Node, error) +func CreateRouter(conf *config.Config, rc *redis.Client, node LocalNode) Router { + if rc != nil { + return NewRedisRouter(node, rc) + } + + // local routing and store + logger.Infow("using single-node routing") + return NewLocalRouter(node) } diff --git a/pkg/routing/localrouter.go b/pkg/routing/localrouter.go index b4d6e2dea..53972c417 100644 --- a/pkg/routing/localrouter.go +++ b/pkg/routing/localrouter.go @@ -42,7 +42,7 @@ func (r *LocalRouter) GetNodeForRoom(ctx context.Context, roomName string) (*liv return node, nil } -func (r *LocalRouter) SetNodeForRoom(ctx context.Context, roomName string, nodeId string) error { +func (r *LocalRouter) SetNodeForRoom(ctx context.Context, roomName, nodeId string) error { return nil } @@ -142,7 +142,7 @@ func (r *LocalRouter) Start() error { return nil } -func (r *LocalRouter) PreStop() { +func (r *LocalRouter) Drain() { r.currentNode.State = livekit.NodeState_SHUTTING_DOWN } diff --git a/pkg/routing/redisrouter.go b/pkg/routing/redisrouter.go index def478c35..0e50ef17f 100644 --- a/pkg/routing/redisrouter.go +++ b/pkg/routing/redisrouter.go @@ -5,13 +5,13 @@ import ( "time" "github.com/go-redis/redis/v8" - "github.com/livekit/livekit-server/pkg/routing/selector" "github.com/livekit/protocol/logger" livekit "github.com/livekit/protocol/proto" "github.com/livekit/protocol/utils" "github.com/pkg/errors" "google.golang.org/protobuf/proto" + "github.com/livekit/livekit-server/pkg/routing/selector" "github.com/livekit/livekit-server/pkg/utils/stats" ) @@ -26,6 +26,7 @@ const ( // Because type RedisRouter struct { LocalRouter + rc *redis.Client ctx context.Context isStarted utils.AtomicFlag @@ -85,7 +86,7 @@ func (r *RedisRouter) GetNodeForRoom(ctx context.Context, roomName string) (*liv return r.GetNode(nodeId) } -func (r *RedisRouter) SetNodeForRoom(ctx context.Context, roomName string, nodeId string) error { +func (r *RedisRouter) SetNodeForRoom(ctx context.Context, roomName, nodeId string) error { return r.rc.HSet(r.ctx, NodeRoomKey, roomName, nodeId).Err() } @@ -259,7 +260,7 @@ func (r *RedisRouter) Start() error { } } -func (r *RedisRouter) PreStop() { +func (r *RedisRouter) Drain() { r.currentNode.State = livekit.NodeState_SHUTTING_DOWN r.RegisterNode() } diff --git a/pkg/routing/routingfakes/fake_node_selector.go b/pkg/routing/routingfakes/fake_node_selector.go deleted file mode 100644 index ddf3ae45f..000000000 --- a/pkg/routing/routingfakes/fake_node_selector.go +++ /dev/null @@ -1,124 +0,0 @@ -// Code generated by counterfeiter. DO NOT EDIT. -package routingfakes - -import ( - "sync" - - "github.com/livekit/livekit-server/pkg/routing" - livekit "github.com/livekit/protocol/proto" -) - -type FakeNodeSelector struct { - SelectNodeStub func([]*livekit.Node, *livekit.Room) (*livekit.Node, error) - selectNodeMutex sync.RWMutex - selectNodeArgsForCall []struct { - arg1 []*livekit.Node - arg2 *livekit.Room - } - selectNodeReturns struct { - result1 *livekit.Node - result2 error - } - selectNodeReturnsOnCall map[int]struct { - result1 *livekit.Node - result2 error - } - invocations map[string][][]interface{} - invocationsMutex sync.RWMutex -} - -func (fake *FakeNodeSelector) SelectNode(arg1 []*livekit.Node, arg2 *livekit.Room) (*livekit.Node, error) { - var arg1Copy []*livekit.Node - if arg1 != nil { - arg1Copy = make([]*livekit.Node, len(arg1)) - copy(arg1Copy, arg1) - } - fake.selectNodeMutex.Lock() - ret, specificReturn := fake.selectNodeReturnsOnCall[len(fake.selectNodeArgsForCall)] - fake.selectNodeArgsForCall = append(fake.selectNodeArgsForCall, struct { - arg1 []*livekit.Node - arg2 *livekit.Room - }{arg1Copy, arg2}) - stub := fake.SelectNodeStub - fakeReturns := fake.selectNodeReturns - fake.recordInvocation("SelectNode", []interface{}{arg1Copy, arg2}) - fake.selectNodeMutex.Unlock() - if stub != nil { - return stub(arg1, arg2) - } - if specificReturn { - return ret.result1, ret.result2 - } - return fakeReturns.result1, fakeReturns.result2 -} - -func (fake *FakeNodeSelector) SelectNodeCallCount() int { - fake.selectNodeMutex.RLock() - defer fake.selectNodeMutex.RUnlock() - return len(fake.selectNodeArgsForCall) -} - -func (fake *FakeNodeSelector) SelectNodeCalls(stub func([]*livekit.Node, *livekit.Room) (*livekit.Node, error)) { - fake.selectNodeMutex.Lock() - defer fake.selectNodeMutex.Unlock() - fake.SelectNodeStub = stub -} - -func (fake *FakeNodeSelector) SelectNodeArgsForCall(i int) ([]*livekit.Node, *livekit.Room) { - fake.selectNodeMutex.RLock() - defer fake.selectNodeMutex.RUnlock() - argsForCall := fake.selectNodeArgsForCall[i] - return argsForCall.arg1, argsForCall.arg2 -} - -func (fake *FakeNodeSelector) SelectNodeReturns(result1 *livekit.Node, result2 error) { - fake.selectNodeMutex.Lock() - defer fake.selectNodeMutex.Unlock() - fake.SelectNodeStub = nil - fake.selectNodeReturns = struct { - result1 *livekit.Node - result2 error - }{result1, result2} -} - -func (fake *FakeNodeSelector) SelectNodeReturnsOnCall(i int, result1 *livekit.Node, result2 error) { - fake.selectNodeMutex.Lock() - defer fake.selectNodeMutex.Unlock() - fake.SelectNodeStub = nil - if fake.selectNodeReturnsOnCall == nil { - fake.selectNodeReturnsOnCall = make(map[int]struct { - result1 *livekit.Node - result2 error - }) - } - fake.selectNodeReturnsOnCall[i] = struct { - result1 *livekit.Node - result2 error - }{result1, result2} -} - -func (fake *FakeNodeSelector) Invocations() map[string][][]interface{} { - fake.invocationsMutex.RLock() - defer fake.invocationsMutex.RUnlock() - fake.selectNodeMutex.RLock() - defer fake.selectNodeMutex.RUnlock() - copiedInvocations := map[string][][]interface{}{} - for key, value := range fake.invocations { - copiedInvocations[key] = value - } - return copiedInvocations -} - -func (fake *FakeNodeSelector) recordInvocation(key string, args []interface{}) { - fake.invocationsMutex.Lock() - defer fake.invocationsMutex.Unlock() - if fake.invocations == nil { - fake.invocations = map[string][][]interface{}{} - } - if fake.invocations[key] == nil { - fake.invocations[key] = [][]interface{}{} - } - fake.invocations[key] = append(fake.invocations[key], args) -} - -var _ routing.NodeSelector = new(FakeNodeSelector) diff --git a/pkg/routing/routingfakes/fake_router.go b/pkg/routing/routingfakes/fake_router.go index 8a5849095..e6183c087 100644 --- a/pkg/routing/routingfakes/fake_router.go +++ b/pkg/routing/routingfakes/fake_router.go @@ -22,6 +22,10 @@ type FakeRouter struct { clearRoomStateReturnsOnCall map[int]struct { result1 error } + DrainStub func() + drainMutex sync.RWMutex + drainArgsForCall []struct { + } GetNodeStub func(string) (*livekit.Node, error) getNodeMutex sync.RWMutex getNodeArgsForCall []struct { @@ -71,10 +75,6 @@ type FakeRouter struct { onRTCMessageArgsForCall []struct { arg1 routing.RTCMessageCallback } - PreStopStub func() - preStopMutex sync.RWMutex - preStopArgsForCall []struct { - } RegisterNodeStub func() error registerNodeMutex sync.RWMutex registerNodeArgsForCall []struct { @@ -231,6 +231,30 @@ func (fake *FakeRouter) ClearRoomStateReturnsOnCall(i int, result1 error) { }{result1} } +func (fake *FakeRouter) Drain() { + fake.drainMutex.Lock() + fake.drainArgsForCall = append(fake.drainArgsForCall, struct { + }{}) + stub := fake.DrainStub + fake.recordInvocation("Drain", []interface{}{}) + fake.drainMutex.Unlock() + if stub != nil { + fake.DrainStub() + } +} + +func (fake *FakeRouter) DrainCallCount() int { + fake.drainMutex.RLock() + defer fake.drainMutex.RUnlock() + return len(fake.drainArgsForCall) +} + +func (fake *FakeRouter) DrainCalls(stub func()) { + fake.drainMutex.Lock() + defer fake.drainMutex.Unlock() + fake.DrainStub = stub +} + func (fake *FakeRouter) GetNode(arg1 string) (*livekit.Node, error) { fake.getNodeMutex.Lock() ret, specificReturn := fake.getNodeReturnsOnCall[len(fake.getNodeArgsForCall)] @@ -480,30 +504,6 @@ func (fake *FakeRouter) OnRTCMessageArgsForCall(i int) routing.RTCMessageCallbac return argsForCall.arg1 } -func (fake *FakeRouter) PreStop() { - fake.preStopMutex.Lock() - fake.preStopArgsForCall = append(fake.preStopArgsForCall, struct { - }{}) - stub := fake.PreStopStub - fake.recordInvocation("PreStop", []interface{}{}) - fake.preStopMutex.Unlock() - if stub != nil { - fake.PreStopStub() - } -} - -func (fake *FakeRouter) PreStopCallCount() int { - fake.preStopMutex.RLock() - defer fake.preStopMutex.RUnlock() - return len(fake.preStopArgsForCall) -} - -func (fake *FakeRouter) PreStopCalls(stub func()) { - fake.preStopMutex.Lock() - defer fake.preStopMutex.Unlock() - fake.PreStopStub = stub -} - func (fake *FakeRouter) RegisterNode() error { fake.registerNodeMutex.Lock() ret, specificReturn := fake.registerNodeReturnsOnCall[len(fake.registerNodeArgsForCall)] @@ -944,6 +944,8 @@ func (fake *FakeRouter) Invocations() map[string][][]interface{} { defer fake.invocationsMutex.RUnlock() fake.clearRoomStateMutex.RLock() defer fake.clearRoomStateMutex.RUnlock() + fake.drainMutex.RLock() + defer fake.drainMutex.RUnlock() fake.getNodeMutex.RLock() defer fake.getNodeMutex.RUnlock() fake.getNodeForRoomMutex.RLock() @@ -954,8 +956,6 @@ func (fake *FakeRouter) Invocations() map[string][][]interface{} { defer fake.onNewParticipantRTCMutex.RUnlock() fake.onRTCMessageMutex.RLock() defer fake.onRTCMessageMutex.RUnlock() - fake.preStopMutex.RLock() - defer fake.preStopMutex.RUnlock() fake.registerNodeMutex.RLock() defer fake.registerNodeMutex.RUnlock() fake.removeDeadNodesMutex.RLock() diff --git a/pkg/routing/selector/interfaces.go b/pkg/routing/selector/interfaces.go new file mode 100644 index 000000000..ffcb47b14 --- /dev/null +++ b/pkg/routing/selector/interfaces.go @@ -0,0 +1,40 @@ +package selector + +import ( + "errors" + + livekit "github.com/livekit/protocol/proto" + + "github.com/livekit/livekit-server/pkg/config" +) + +var ErrUnsupportedSelector = errors.New("unsupported node selector") + +// NodeSelector selects an appropriate node to run the current session +type NodeSelector interface { + SelectNode(nodes []*livekit.Node) (*livekit.Node, error) +} + +func CreateNodeSelector(conf *config.Config) (NodeSelector, error) { + kind := conf.NodeSelector.Kind + if kind == "" { + kind = "random" + } + switch kind { + case "sysload": + return &SystemLoadSelector{ + SysloadLimit: conf.NodeSelector.SysloadLimit, + }, nil + case "regionaware": + s, err := NewRegionAwareSelector(conf.Region, conf.NodeSelector.Regions) + if err != nil { + return nil, err + } + s.SysloadLimit = conf.NodeSelector.SysloadLimit + return s, nil + case "random": + return &RandomSelector{}, nil + default: + return nil, ErrUnsupportedSelector + } +} diff --git a/pkg/routing/selector/random.go b/pkg/routing/selector/random.go index d67240f8f..67d4f045f 100644 --- a/pkg/routing/selector/random.go +++ b/pkg/routing/selector/random.go @@ -9,7 +9,7 @@ import ( type RandomSelector struct { } -func (s *RandomSelector) SelectNode(nodes []*livekit.Node, room *livekit.Room) (*livekit.Node, error) { +func (s *RandomSelector) SelectNode(nodes []*livekit.Node) (*livekit.Node, error) { nodes = GetAvailableNodes(nodes) if len(nodes) == 0 { return nil, ErrNoAvailableNodes diff --git a/pkg/routing/selector/regionaware.go b/pkg/routing/selector/regionaware.go index 97680501a..06556a317 100644 --- a/pkg/routing/selector/regionaware.go +++ b/pkg/routing/selector/regionaware.go @@ -5,25 +5,19 @@ import ( livekit "github.com/livekit/protocol/proto" "github.com/thoas/go-funk" -) -// RegionConfig lists available regions and their latitude/longitude, so the selector would prefer -// regions that are closer -type RegionConfig struct { - Name string `yaml:"name"` - Lat float64 `yaml:"lat"` - Lon float64 `yaml:"lon"` -} + "github.com/livekit/livekit-server/pkg/config" +) // RegionAwareSelector prefers available nodes that are closest to the region of the current instance type RegionAwareSelector struct { SystemLoadSelector CurrentRegion string regionDistances map[string]float64 - regions []RegionConfig + regions []config.RegionConfig } -func NewRegionAwareSelector(currentRegion string, regions []RegionConfig) (*RegionAwareSelector, error) { +func NewRegionAwareSelector(currentRegion string, regions []config.RegionConfig) (*RegionAwareSelector, error) { if currentRegion == "" { return nil, ErrCurrentRegionNotSet } @@ -34,7 +28,7 @@ func NewRegionAwareSelector(currentRegion string, regions []RegionConfig) (*Regi regions: regions, } - var currentRC *RegionConfig + var currentRC *config.RegionConfig for _, region := range regions { if region.Name == currentRegion { @@ -56,7 +50,7 @@ func NewRegionAwareSelector(currentRegion string, regions []RegionConfig) (*Regi return s, nil } -func (s *RegionAwareSelector) SelectNode(nodes []*livekit.Node, room *livekit.Room) (*livekit.Node, error) { +func (s *RegionAwareSelector) SelectNode(nodes []*livekit.Node) (*livekit.Node, error) { nodes, err := s.SystemLoadSelector.filterNodes(nodes) if err != nil { return nil, err diff --git a/pkg/routing/selector/regionaware_test.go b/pkg/routing/selector/regionaware_test.go index 29a92b476..19a0bfbfb 100644 --- a/pkg/routing/selector/regionaware_test.go +++ b/pkg/routing/selector/regionaware_test.go @@ -4,10 +4,12 @@ import ( "testing" "time" - "github.com/livekit/livekit-server/pkg/routing/selector" livekit "github.com/livekit/protocol/proto" "github.com/livekit/protocol/utils" "github.com/stretchr/testify/require" + + "github.com/livekit/livekit-server/pkg/config" + "github.com/livekit/livekit-server/pkg/routing/selector" ) const ( @@ -18,7 +20,7 @@ const ( ) func TestRegionAwareRouting(t *testing.T) { - rc := []selector.RegionConfig{ + rc := []config.RegionConfig{ { Name: regionWest, Lat: 37.64046607830567, @@ -42,7 +44,7 @@ func TestRegionAwareRouting(t *testing.T) { s, err := selector.NewRegionAwareSelector(regionEast, nil) require.NoError(t, err) - node, err := s.SelectNode(nodes, nil) + node, err := s.SelectNode(nodes) require.NoError(t, err) require.NotNil(t, node) }) @@ -59,7 +61,7 @@ func TestRegionAwareRouting(t *testing.T) { require.NoError(t, err) s.SysloadLimit = loadLimit - node, err := s.SelectNode(nodes, nil) + node, err := s.SelectNode(nodes) require.NoError(t, err) require.Equal(t, expectedNode, node) }) @@ -76,7 +78,7 @@ func TestRegionAwareRouting(t *testing.T) { require.NoError(t, err) s.SysloadLimit = loadLimit - node, err := s.SelectNode(nodes, nil) + node, err := s.SelectNode(nodes) require.NoError(t, err) require.Equal(t, expectedNode, node) }) @@ -92,7 +94,7 @@ func TestRegionAwareRouting(t *testing.T) { require.NoError(t, err) s.SysloadLimit = loadLimit - node, err := s.SelectNode(nodes, nil) + node, err := s.SelectNode(nodes) require.NoError(t, err) require.Equal(t, expectedNode, node) }) @@ -110,7 +112,7 @@ func TestRegionAwareRouting(t *testing.T) { require.NoError(t, err) s.SysloadLimit = loadLimit - node, err := s.SelectNode(nodes, nil) + node, err := s.SelectNode(nodes) require.NoError(t, err) require.Equal(t, expectedNode, node) }) @@ -122,7 +124,7 @@ func TestRegionAwareRouting(t *testing.T) { s, err := selector.NewRegionAwareSelector(regionEast, rc) require.NoError(t, err) - node, err := s.SelectNode(nodes, nil) + node, err := s.SelectNode(nodes) require.NoError(t, err) require.NotNil(t, node) }) diff --git a/pkg/routing/selector/sysload.go b/pkg/routing/selector/sysload.go index 22a2af95a..2bc41182a 100644 --- a/pkg/routing/selector/sysload.go +++ b/pkg/routing/selector/sysload.go @@ -33,7 +33,7 @@ func (s *SystemLoadSelector) filterNodes(nodes []*livekit.Node) ([]*livekit.Node return nodes, nil } -func (s *SystemLoadSelector) SelectNode(nodes []*livekit.Node, room *livekit.Room) (*livekit.Node, error) { +func (s *SystemLoadSelector) SelectNode(nodes []*livekit.Node) (*livekit.Node, error) { nodes, err := s.filterNodes(nodes) if err != nil { return nil, err diff --git a/pkg/routing/selector/sysload_test.go b/pkg/routing/selector/sysload_test.go index 7667128bd..8af8ea4ef 100644 --- a/pkg/routing/selector/sysload_test.go +++ b/pkg/routing/selector/sysload_test.go @@ -4,9 +4,10 @@ import ( "testing" "time" - "github.com/livekit/livekit-server/pkg/routing/selector" livekit "github.com/livekit/protocol/proto" "github.com/stretchr/testify/require" + + "github.com/livekit/livekit-server/pkg/routing/selector" ) var ( @@ -33,19 +34,19 @@ func TestSystemLoadSelector_SelectNode(t *testing.T) { selector := selector.SystemLoadSelector{SysloadLimit: 1.0} nodes := []*livekit.Node{} - _, err := selector.SelectNode(nodes, nil) + _, err := selector.SelectNode(nodes) require.Error(t, err, "should error no available nodes") // Select a node with high load when no nodes with low load are available nodes = []*livekit.Node{nodeLoadHigh} - if _, err := selector.SelectNode(nodes, nil); err != nil { + if _, err := selector.SelectNode(nodes); err != nil { t.Error(err) } // Select a node with low load when available nodes = []*livekit.Node{nodeLoadLow, nodeLoadHigh} for i := 0; i < 5; i++ { - node, err := selector.SelectNode(nodes, nil) + node, err := selector.SelectNode(nodes) if err != nil { t.Error(err) } diff --git a/pkg/routing/selector/utils.go b/pkg/routing/selector/utils.go index 7e77cad17..85db8efb6 100644 --- a/pkg/routing/selector/utils.go +++ b/pkg/routing/selector/utils.go @@ -7,9 +7,7 @@ import ( "github.com/thoas/go-funk" ) -const ( - AvailableSeconds = 5 -) +const AvailableSeconds = 5 // checks if a node has been updated recently to be considered for selection func IsAvailable(node *livekit.Node) bool { diff --git a/pkg/service/errors.go b/pkg/service/errors.go index 2d6c79e9b..992e45549 100644 --- a/pkg/service/errors.go +++ b/pkg/service/errors.go @@ -9,5 +9,4 @@ var ( ErrParticipantNotFound = errors.New("participant does not exist") ErrTrackNotFound = errors.New("track is not found") ErrWebHookMissingAPIKey = errors.New("api_key is required to use webhooks") - ErrUnsupportedSelector = errors.New("unsupported node selector") ) diff --git a/pkg/service/interfaces.go b/pkg/service/interfaces.go index 2aab7ff00..4de1e3590 100644 --- a/pkg/service/interfaces.go +++ b/pkg/service/interfaces.go @@ -13,7 +13,6 @@ import ( //go:generate go run github.com/maxbrunsfeld/counterfeiter/v6 -generate // encapsulates CRUD operations for room settings -// look up participant //counterfeiter:generate . RoomStore type RoomStore interface { StoreRoom(ctx context.Context, room *livekit.Room) error diff --git a/pkg/service/roomallocator.go b/pkg/service/roomallocator.go index 11353eca1..e55bc7cd6 100644 --- a/pkg/service/roomallocator.go +++ b/pkg/service/roomallocator.go @@ -4,29 +4,34 @@ import ( "context" "time" - "github.com/livekit/livekit-server/pkg/routing/selector" "github.com/livekit/protocol/logger" livekit "github.com/livekit/protocol/proto" "github.com/livekit/protocol/utils" "github.com/livekit/livekit-server/pkg/config" "github.com/livekit/livekit-server/pkg/routing" + "github.com/livekit/livekit-server/pkg/routing/selector" ) type RoomAllocator struct { config *config.Config router routing.Router - selector routing.NodeSelector + selector selector.NodeSelector roomStore RoomStore } -func NewRoomAllocator(conf *config.Config, router routing.Router, selector routing.NodeSelector, rs RoomStore) *RoomAllocator { +func NewRoomAllocator(conf *config.Config, router routing.Router, rs RoomStore) (*RoomAllocator, error) { + ns, err := selector.CreateNodeSelector(conf) + if err != nil { + return nil, err + } + return &RoomAllocator{ config: conf, router: router, - selector: selector, + selector: ns, roomStore: rs, - } + }, nil } // CreateRoom creates a new room from a request and allocates it to a node to handle @@ -64,35 +69,36 @@ func (r *RoomAllocator) CreateRoom(ctx context.Context, req *livekit.CreateRoomR return nil, err } - // Is that node still available? - node, err := r.router.GetNodeForRoom(ctx, rm.Name) + // check if room already assigned + existing, err := r.router.GetNodeForRoom(ctx, rm.Name) if err != routing.ErrNotFound && err != nil { return nil, err } - // keep it on that node - if err == nil && selector.IsAvailable(node) { + // if already assigned and still available, keep it on that node + if err == nil && selector.IsAvailable(existing) { return rm, nil } // select a new node nodeId := req.NodeId if nodeId == "" { - // select a node for room nodes, err := r.router.ListNodes() if err != nil { return nil, err } - node, err := r.selector.SelectNode(nodes, rm) + node, err := r.selector.SelectNode(nodes) if err != nil { return nil, err } + nodeId = node.Id } logger.Debugw("selected node for room", "room", rm.Name, "roomID", rm.Sid, "nodeID", nodeId) - if err := r.router.SetNodeForRoom(ctx, req.Name, nodeId); err != nil { + err = r.router.SetNodeForRoom(ctx, rm.Name, nodeId) + if err != nil { return nil, err } diff --git a/pkg/service/roomallocator_test.go b/pkg/service/roomallocator_test.go index 76e17e540..2191f59d6 100644 --- a/pkg/service/roomallocator_test.go +++ b/pkg/service/roomallocator_test.go @@ -4,7 +4,6 @@ import ( "context" "testing" - "github.com/livekit/livekit-server/pkg/routing/selector" livekit "github.com/livekit/protocol/proto" "github.com/stretchr/testify/require" @@ -32,12 +31,12 @@ func newTestRoomAllocator(t *testing.T) (*service.RoomAllocator, *config.Config) router := &routingfakes.FakeRouter{} conf, err := config.NewConfig("", nil) require.NoError(t, err) - selector := &selector.RandomSelector{} node, err := routing.NewLocalNode(conf) require.NoError(t, err) router.GetNodeForRoomReturns(node, nil) - ra := service.NewRoomAllocator(conf, router, selector, store) + ra, err := service.NewRoomAllocator(conf, router, store) + require.NoError(t, err) return ra, conf } diff --git a/pkg/service/roommanager.go b/pkg/service/roommanager.go index 32bb96973..eb3dac2bc 100644 --- a/pkg/service/roommanager.go +++ b/pkg/service/roommanager.go @@ -27,7 +27,6 @@ type LocalRoomManager struct { RoomStore lock sync.RWMutex - selector routing.NodeSelector router routing.Router currentNode routing.LocalNode notifier webhook.Notifier @@ -37,20 +36,20 @@ type LocalRoomManager struct { rooms map[string]*rtc.Room } -func NewLocalRoomManager(rp RoomStore, router routing.Router, currentNode routing.LocalNode, selector routing.NodeSelector, - notifier webhook.Notifier, conf *config.Config) (*LocalRoomManager, error) { +func NewLocalRoomManager(conf *config.Config, rs RoomStore, router routing.Router, currentNode routing.LocalNode, + notifier webhook.Notifier) (*LocalRoomManager, error) { + rtcConf, err := rtc.NewWebRTCConfig(conf, currentNode.Ip) if err != nil { return nil, err } r := &LocalRoomManager{ - RoomStore: rp, + RoomStore: rs, lock: sync.RWMutex{}, rtcConfig: rtcConf, config: conf, router: router, - selector: selector, notifier: notifier, currentNode: currentNode, webhookPool: workerpool.New(1), diff --git a/pkg/service/roomservice.go b/pkg/service/roomservice.go index 6aa62b585..8be3744e1 100644 --- a/pkg/service/roomservice.go +++ b/pkg/service/roomservice.go @@ -14,7 +14,6 @@ import ( // A rooms service that supports a single node type RoomService struct { router routing.Router - selector routing.NodeSelector roomAllocator *RoomAllocator roomStore RoomStore } diff --git a/pkg/service/server.go b/pkg/service/server.go index c92839fce..4641df76b 100644 --- a/pkg/service/server.go +++ b/pkg/service/server.go @@ -207,7 +207,7 @@ func (s *LivekitServer) Start() error { func (s *LivekitServer) Stop(force bool) { // wait for all participants to exit - s.router.PreStop() + s.router.Drain() partTicker := time.NewTicker(5 * time.Second) waitingForParticipants := !force && s.roomManager.HasParticipants() for waitingForParticipants { diff --git a/pkg/service/utils.go b/pkg/service/utils.go index 74c163c74..001ece188 100644 --- a/pkg/service/utils.go +++ b/pkg/service/utils.go @@ -1,153 +1,14 @@ package service import ( - "context" - "fmt" "net/http" - "os" "regexp" - "github.com/go-redis/redis/v8" - "github.com/google/wire" "github.com/livekit/protocol/auth" "github.com/livekit/protocol/logger" livekit "github.com/livekit/protocol/proto" - "github.com/livekit/protocol/utils" - "github.com/livekit/protocol/webhook" - "github.com/pkg/errors" - - "github.com/livekit/livekit-server/pkg/config" - "github.com/livekit/livekit-server/pkg/routing" - "github.com/livekit/livekit-server/pkg/routing/selector" ) -var ServiceSet = wire.NewSet( - createRedisClient, - createMessageBus, - createRouter, - createStore, - CreateKeyProvider, - CreateWebhookNotifier, - CreateNodeSelector, - NewRecordingService, - NewRoomAllocator, - NewRoomService, - NewRTCService, - NewLivekitServer, - NewLocalRoomManager, - newTurnAuthHandler, - NewTurnServer, - config.GetAudioConfig, - wire.Bind(new(RoomManager), new(*LocalRoomManager)), - wire.Bind(new(livekit.RoomService), new(*RoomService)), -) - -func CreateKeyProvider(conf *config.Config) (auth.KeyProvider, error) { - // prefer keyfile if set - if conf.KeyFile != "" { - if st, err := os.Stat(conf.KeyFile); err != nil { - return nil, err - } else if st.Mode().Perm() != 0600 { - return nil, fmt.Errorf("key file must have permission set to 600") - } - f, err := os.Open(conf.KeyFile) - if err != nil { - return nil, err - } - defer func() { - _ = f.Close() - }() - return auth.NewFileBasedKeyProviderFromReader(f) - } - - if len(conf.Keys) == 0 { - return nil, errors.New("one of key-file or keys must be provided in order to support a secure installation") - } - - return auth.NewFileBasedKeyProviderFromMap(conf.Keys), nil -} - -func CreateWebhookNotifier(conf *config.Config, provider auth.KeyProvider) (webhook.Notifier, error) { - wc := conf.WebHook - if len(wc.URLs) == 0 { - return nil, nil - } - secret := provider.GetSecret(wc.APIKey) - if secret == "" { - return nil, ErrWebHookMissingAPIKey - } - - return webhook.NewNotifier(wc.APIKey, secret, wc.URLs), nil -} - -func CreateNodeSelector(conf *config.Config) (routing.NodeSelector, error) { - kind := conf.NodeSelector.Kind - if kind == "" { - kind = "random" - } - switch kind { - case "sysload": - return &selector.SystemLoadSelector{ - SysloadLimit: conf.NodeSelector.SysloadLimit, - }, nil - case "regionaware": - s, err := selector.NewRegionAwareSelector(conf.Region, conf.NodeSelector.Regions) - if err != nil { - return nil, err - } - s.SysloadLimit = conf.NodeSelector.SysloadLimit - return s, nil - case "random": - return &selector.RandomSelector{}, nil - default: - return nil, ErrUnsupportedSelector - } -} - -func createRedisClient(conf *config.Config) (*redis.Client, error) { - if !conf.HasRedis() { - return nil, nil - } - - logger.Infow("using multi-node routing via redis", "addr", conf.Redis.Address) - rc := redis.NewClient(&redis.Options{ - Addr: conf.Redis.Address, - Username: conf.Redis.Username, - Password: conf.Redis.Password, - DB: conf.Redis.DB, - }) - if err := rc.Ping(context.Background()).Err(); err != nil { - err = errors.Wrap(err, "unable to connect to redis") - return nil, err - } - - return rc, nil -} - -func createMessageBus(rc *redis.Client) utils.MessageBus { - if rc == nil { - return nil - } - return utils.NewRedisMessageBus(rc) -} - -func createRouter(rc *redis.Client, node routing.LocalNode) routing.Router { - if rc != nil { - return routing.NewRedisRouter(node, rc) - } - - // local routing and store - logger.Infow("using single-node routing") - return routing.NewLocalRouter(node) -} - -func createStore(rc *redis.Client) RoomStore { - if rc != nil { - return NewRedisRoomStore(rc) - } - return NewLocalRoomStore() -} - func handleError(w http.ResponseWriter, status int, msg string) { // GetLogger already with extra depth 1 logger.GetLogger().V(1).Info("error handling request", "error", msg, "status", status) diff --git a/pkg/service/wire.go b/pkg/service/wire.go index 48313d4a5..14ae55d8c 100644 --- a/pkg/service/wire.go +++ b/pkg/service/wire.go @@ -3,7 +3,19 @@ package service import ( + "context" + "fmt" + "os" + + "github.com/go-redis/redis/v8" "github.com/google/wire" + "github.com/pkg/errors" + + "github.com/livekit/protocol/auth" + "github.com/livekit/protocol/logger" + livekit "github.com/livekit/protocol/proto" + "github.com/livekit/protocol/utils" + "github.com/livekit/protocol/webhook" "github.com/livekit/livekit-server/pkg/config" "github.com/livekit/livekit-server/pkg/routing" @@ -11,18 +23,102 @@ import ( func InitializeServer(conf *config.Config, currentNode routing.LocalNode) (*LivekitServer, error) { wire.Build( - ServiceSet, + createRedisClient, + createMessageBus, + createStore, + createKeyProvider, + createWebhookNotifier, + routing.CreateRouter, + NewRecordingService, + NewRoomAllocator, + NewRoomService, + NewRTCService, + NewLocalRoomManager, + newTurnAuthHandler, + NewTurnServer, + wire.Bind(new(livekit.RoomService), new(*RoomService)), + NewLivekitServer, ) return &LivekitServer{}, nil } func InitializeRouter(conf *config.Config, currentNode routing.LocalNode) (routing.Router, error) { wire.Build( - wire.NewSet( - createRedisClient, - createRouter, - ), + createRedisClient, + routing.CreateRouter, ) return nil, nil } + +func createKeyProvider(conf *config.Config) (auth.KeyProvider, error) { + // prefer keyfile if set + if conf.KeyFile != "" { + if st, err := os.Stat(conf.KeyFile); err != nil { + return nil, err + } else if st.Mode().Perm() != 0600 { + return nil, fmt.Errorf("key file must have permission set to 600") + } + f, err := os.Open(conf.KeyFile) + if err != nil { + return nil, err + } + defer func() { + _ = f.Close() + }() + return auth.NewFileBasedKeyProviderFromReader(f) + } + + if len(conf.Keys) == 0 { + return nil, errors.New("one of key-file or keys must be provided in order to support a secure installation") + } + + return auth.NewFileBasedKeyProviderFromMap(conf.Keys), nil +} + +func createWebhookNotifier(conf *config.Config, provider auth.KeyProvider) (webhook.Notifier, error) { + wc := conf.WebHook + if len(wc.URLs) == 0 { + return nil, nil + } + secret := provider.GetSecret(wc.APIKey) + if secret == "" { + return nil, ErrWebHookMissingAPIKey + } + + return webhook.NewNotifier(wc.APIKey, secret, wc.URLs), nil +} + +func createRedisClient(conf *config.Config) (*redis.Client, error) { + if !conf.HasRedis() { + return nil, nil + } + + logger.Infow("using multi-node routing via redis", "addr", conf.Redis.Address) + rc := redis.NewClient(&redis.Options{ + Addr: conf.Redis.Address, + Username: conf.Redis.Username, + Password: conf.Redis.Password, + DB: conf.Redis.DB, + }) + if err := rc.Ping(context.Background()).Err(); err != nil { + err = errors.Wrap(err, "unable to connect to redis") + return nil, err + } + + return rc, nil +} + +func createMessageBus(rc *redis.Client) utils.MessageBus { + if rc == nil { + return nil + } + return utils.NewRedisMessageBus(rc) +} + +func createStore(rc *redis.Client) RoomStore { + if rc != nil { + return NewRedisRoomStore(rc) + } + return NewLocalRoomStore() +} diff --git a/pkg/service/wire_gen.go b/pkg/service/wire_gen.go index 2f6762f94..f7d3557a8 100644 --- a/pkg/service/wire_gen.go +++ b/pkg/service/wire_gen.go @@ -6,8 +6,17 @@ package service import ( + "context" + "fmt" + "github.com/go-redis/redis/v8" "github.com/livekit/livekit-server/pkg/config" "github.com/livekit/livekit-server/pkg/routing" + "github.com/livekit/protocol/auth" + "github.com/livekit/protocol/logger" + "github.com/livekit/protocol/utils" + "github.com/livekit/protocol/webhook" + "github.com/pkg/errors" + "os" ) // Injectors from wire.go: @@ -17,29 +26,28 @@ func InitializeServer(conf *config.Config, currentNode routing.LocalNode) (*Live if err != nil { return nil, err } - router := createRouter(client, currentNode) - nodeSelector, err := CreateNodeSelector(conf) + router := routing.CreateRouter(conf, client, currentNode) + roomStore := createStore(client) + roomAllocator, err := NewRoomAllocator(conf, router, roomStore) if err != nil { return nil, err } - roomStore := createStore(client) - roomAllocator := NewRoomAllocator(conf, router, nodeSelector, roomStore) roomService, err := NewRoomService(roomAllocator, roomStore, router) if err != nil { return nil, err } messageBus := createMessageBus(client) - keyProvider, err := CreateKeyProvider(conf) + keyProvider, err := createKeyProvider(conf) if err != nil { return nil, err } - notifier, err := CreateWebhookNotifier(conf, keyProvider) + notifier, err := createWebhookNotifier(conf, keyProvider) if err != nil { return nil, err } recordingService := NewRecordingService(messageBus, notifier) rtcService := NewRTCService(conf, roomAllocator, router, currentNode) - localRoomManager, err := NewLocalRoomManager(roomStore, router, currentNode, nodeSelector, notifier, conf) + localRoomManager, err := NewLocalRoomManager(conf, roomStore, router, currentNode, notifier) if err != nil { return nil, err } @@ -60,6 +68,79 @@ func InitializeRouter(conf *config.Config, currentNode routing.LocalNode) (routi if err != nil { return nil, err } - router := createRouter(client, currentNode) + router := routing.CreateRouter(conf, client, currentNode) return router, nil } + +// wire.go: + +func createKeyProvider(conf *config.Config) (auth.KeyProvider, error) { + + if conf.KeyFile != "" { + if st, err := os.Stat(conf.KeyFile); err != nil { + return nil, err + } else if st.Mode().Perm() != 0600 { + return nil, fmt.Errorf("key file must have permission set to 600") + } + f, err := os.Open(conf.KeyFile) + if err != nil { + return nil, err + } + defer func() { + _ = f.Close() + }() + return auth.NewFileBasedKeyProviderFromReader(f) + } + + if len(conf.Keys) == 0 { + return nil, errors.New("one of key-file or keys must be provided in order to support a secure installation") + } + + return auth.NewFileBasedKeyProviderFromMap(conf.Keys), nil +} + +func createWebhookNotifier(conf *config.Config, provider auth.KeyProvider) (webhook.Notifier, error) { + wc := conf.WebHook + if len(wc.URLs) == 0 { + return nil, nil + } + secret := provider.GetSecret(wc.APIKey) + if secret == "" { + return nil, ErrWebHookMissingAPIKey + } + + return webhook.NewNotifier(wc.APIKey, secret, wc.URLs), nil +} + +func createRedisClient(conf *config.Config) (*redis.Client, error) { + if !conf.HasRedis() { + return nil, nil + } + logger.Infow("using multi-node routing via redis", "addr", conf.Redis.Address) + rc := redis.NewClient(&redis.Options{ + Addr: conf.Redis.Address, + Username: conf.Redis.Username, + Password: conf.Redis.Password, + DB: conf.Redis.DB, + }) + if err := rc.Ping(context.Background()).Err(); err != nil { + err = errors.Wrap(err, "unable to connect to redis") + return nil, err + } + + return rc, nil +} + +func createMessageBus(rc *redis.Client) utils.MessageBus { + if rc == nil { + return nil + } + return utils.NewRedisMessageBus(rc) +} + +func createStore(rc *redis.Client) RoomStore { + if rc != nil { + return NewRedisRoomStore(rc) + } + return NewLocalRoomStore() +}