Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions pkg/limits/frontend/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,4 +10,5 @@ type limitsClient interface {
// ExceedsLimits checks if the streams in the request have exceeded their
// per-partition limits.
ExceedsLimits(context.Context, *proto.ExceedsLimitsRequest) ([]*proto.ExceedsLimitsResponse, error)
UpdateRates(context.Context, *proto.UpdateRatesRequest) ([]*proto.UpdateRatesResponse, error)
}
22 changes: 22 additions & 0 deletions pkg/limits/frontend/frontend.go
Original file line number Diff line number Diff line change
Expand Up @@ -137,6 +137,28 @@ func (f *Frontend) ExceedsLimits(ctx context.Context, req *proto.ExceedsLimitsRe
return &proto.ExceedsLimitsResponse{Results: results}, nil
}

func (f *Frontend) UpdateRates(
ctx context.Context,
req *proto.UpdateRatesRequest,
) (*proto.UpdateRatesResponse, error) {
results := make([]*proto.UpdateRatesResult, 0, len(req.Streams))
resps, err := f.limitsClient.UpdateRates(ctx, req)
if err != nil {
// If the entire call failed, then all streams failed.
for _, stream := range req.Streams {
results = append(results, &proto.UpdateRatesResult{
StreamHash: stream.StreamHash,
Rate: 0,
})
}
} else {
for _, resp := range resps {
results = append(results, resp.Results...)
}
}
return &proto.UpdateRatesResponse{Results: results}, nil
}

func (f *Frontend) CheckReady(ctx context.Context) error {
if f.State() != services.Running {
return fmt.Errorf("service is not running: %v", f.State())
Expand Down
5 changes: 5 additions & 0 deletions pkg/limits/frontend/mock_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,11 @@ func (m *mockLimitsClient) ExceedsLimits(_ context.Context, req *proto.ExceedsLi
return m.exceedsLimitsResponses, m.err
}

func (m *mockLimitsClient) UpdateRates(_ context.Context, _ *proto.UpdateRatesRequest) ([]*proto.UpdateRatesResponse, error) {
// TODO(grobinson): Implement this method.
return nil, nil
}

// mockLimitsProtoClient mocks proto.IngestLimitsClient.
type mockLimitsProtoClient struct {
proto.IngestLimitsClient
Expand Down
276 changes: 197 additions & 79 deletions pkg/limits/frontend/ring.go
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ const (
)

var (
LimitsRead = ring.NewOp([]ring.InstanceState{ring.ACTIVE}, nil)
limitsRead = ring.NewOp([]ring.InstanceState{ring.ACTIVE}, nil)

// defaultZoneCmp compares two zones using [strings.Compare].
defaultZoneCmp = func(a, b string) int {
Expand Down Expand Up @@ -73,18 +73,185 @@ func newRingLimitsClient(
}

// ExceedsLimits implements the [exceedsLimitsGatherer] interface.
func (r *ringLimitsClient) ExceedsLimits(ctx context.Context, req *proto.ExceedsLimitsRequest) ([]*proto.ExceedsLimitsResponse, error) {
func (r *ringLimitsClient) ExceedsLimits(
ctx context.Context,
req *proto.ExceedsLimitsRequest,
) ([]*proto.ExceedsLimitsResponse, error) {
if len(req.Streams) == 0 {
return nil, nil
}
responses := make([]*proto.ExceedsLimitsResponse, 0)
doExceedsLimitsRPC := r.newExceedsLimitsRPCsFunc(&responses)
unanswered, err := r.exhaustAllZones(ctx, req.Tenant, req.Streams, doExceedsLimitsRPC)
if err != nil {
return nil, err
}
// Any unanswered streams after exhausting all zones must be failed.
if len(unanswered) > 0 {
failed := make([]*proto.ExceedsLimitsResult, 0, len(unanswered))
for _, stream := range unanswered {
failed = append(failed, &proto.ExceedsLimitsResult{
StreamHash: stream.StreamHash,
Reason: uint32(limits.ReasonFailed),
})
}
responses = append(responses, &proto.ExceedsLimitsResponse{Results: failed})
}
return responses, nil
}

// UpdateRates implements the [exceedsLimitsGatherer] interface.
func (r *ringLimitsClient) UpdateRates(
ctx context.Context,
req *proto.UpdateRatesRequest,
) ([]*proto.UpdateRatesResponse, error) {
if len(req.Streams) == 0 {
return nil, nil
}
responses := make([]*proto.UpdateRatesResponse, 0)
doUpdateRatesRPC := r.newUpdateRatesRPCsFunc(&responses)
unanswered, err := r.exhaustAllZones(ctx, req.Tenant, req.Streams, doUpdateRatesRPC)
if err != nil {
return nil, err
}
// Any unanswered streams after exhausting all zones have a rate of 0.
if len(unanswered) > 0 {
failed := make([]*proto.UpdateRatesResult, 0, len(unanswered))
for _, stream := range unanswered {
failed = append(failed, &proto.UpdateRatesResult{

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I decided not to leave failed streams out of the response as I think its better we make failed updates explicit in future rather than implicit (i.e. absent from the response) like we do with ExceedsLimits.

StreamHash: stream.StreamHash,
Rate: 0,
})
}
responses = append(responses, &proto.UpdateRatesResponse{Results: failed})
}
return responses, nil
}

// newExceedsLimitsRPCsFunc returns a doRPCsFunc that executes the ExceedsLimits
// RPCs for the instances in a zone.
func (r *ringLimitsClient) newExceedsLimitsRPCsFunc(responses *[]*proto.ExceedsLimitsResponse) doRPCsFunc {
return func(
ctx context.Context,
tenant string,
streams []*proto.StreamMetadata,
zone string,
consumers map[int32]string,
) ([]uint64, error) {
errg, ctx := errgroup.WithContext(ctx)
instancesForStreams := r.instancesForStreams(streams, zone, consumers)
responseCh := make(chan *proto.ExceedsLimitsResponse, len(instancesForStreams))
answeredCh := make(chan uint64, len(streams))
for addr, streams := range instancesForStreams {
errg.Go(func() error {
client, err := r.pool.GetClientFor(addr)
if err != nil {
level.Error(r.logger).Log("msg", "failed to get client for instance", "instance", addr, "err", err.Error())
return nil
}
resp, err := client.(proto.IngestLimitsClient).ExceedsLimits(ctx, &proto.ExceedsLimitsRequest{
Tenant: tenant,
Streams: streams,
})
if err != nil {
level.Error(r.logger).Log("failed check execeed limits for instance", "instance", addr, "err", err.Error())
return nil
}
responseCh <- resp
for _, stream := range streams {
answeredCh <- stream.StreamHash
}
return nil
})
}
_ = errg.Wait()
close(responseCh)
close(answeredCh)
for r := range responseCh {
*responses = append(*responses, r)
}
answered := make([]uint64, 0, len(streams))
for streamHash := range answeredCh {
answered = append(answered, streamHash)
}
return answered, nil
}
}

// newUpdateRatesRPCsFunc returns a doRPCsFunc that executes the UpdateRates
// RPCs for the instances in a zone.
func (r *ringLimitsClient) newUpdateRatesRPCsFunc(responses *[]*proto.UpdateRatesResponse) doRPCsFunc {
return func(
ctx context.Context,
tenant string,
streams []*proto.StreamMetadata,
zone string,
consumers map[int32]string,
) ([]uint64, error) {
errg, ctx := errgroup.WithContext(ctx)
instancesForStreams := r.instancesForStreams(streams, zone, consumers)
responseCh := make(chan *proto.UpdateRatesResponse, len(instancesForStreams))
answeredCh := make(chan uint64, len(streams))
for addr, streams := range instancesForStreams {
errg.Go(func() error {
client, err := r.pool.GetClientFor(addr)
if err != nil {
level.Error(r.logger).Log("msg", "failed to get client for instance", "instance", addr, "err", err.Error())
return nil
}
resp, err := client.(proto.IngestLimitsClient).UpdateRates(ctx, &proto.UpdateRatesRequest{
Tenant: tenant,
Streams: streams,
})
if err != nil {
level.Error(r.logger).Log("failed check execeed limits for instance", "instance", addr, "err", err.Error())
return nil
}
responseCh <- resp
for _, stream := range streams {
answeredCh <- stream.StreamHash
}
return nil
})
}
_ = errg.Wait()
close(responseCh)
close(answeredCh)
for r := range responseCh {
*responses = append(*responses, r)
}
answered := make([]uint64, 0, len(streams))
for streamHash := range answeredCh {
answered = append(answered, streamHash)
}
return answered, nil
}
}

type doRPCsFunc func(
ctx context.Context,
tenant string,
streams []*proto.StreamMetadata,
zone string,
consumers map[int32]string,
) ([]uint64, error)

// exhaustAllZones queries all zones, one at a time, until either all streams
// have been answered or all zones have been exhausted.
func (r *ringLimitsClient) exhaustAllZones(
ctx context.Context,
tenant string,
streams []*proto.StreamMetadata,
doRPCs doRPCsFunc,
) ([]*proto.StreamMetadata, error) {
zonesIter, err := r.allZones(ctx)
if err != nil {
return nil, err
}
// Make a copy of the streams from the request. We will prune this slice
// each time we receive the responses from a zone.
streams := make([]*proto.StreamMetadata, 0, len(req.Streams))
streams = append(streams, req.Streams...)
// Make a copy of the stream. We will prune this slice each time we receive
// the responses from a zone.
unanswered := make([]*proto.StreamMetadata, 0, len(streams))
unanswered = append(unanswered, streams...)
// Query each zone as ordered in zonesToQuery. If a zone answers all
// streams, the request is satisfied and there is no need to query
// subsequent zones. If a zone answers just a subset of streams
Expand All @@ -93,92 +260,28 @@ func (r *ringLimitsClient) ExceedsLimits(ctx context.Context, req *proto.Exceeds
// then query the next zone for the remaining streams. We repeat this
// process until all streams have been queried or we have exhausted all
// zones.
responses := make([]*proto.ExceedsLimitsResponse, 0)
for zone, consumers := range zonesIter {
// All streams been checked against per-tenant limits.
if len(streams) == 0 {
// All streams been answered.
if len(unanswered) == 0 {
break
}
resps, answered, err := r.doExceedsLimitsRPCs(ctx, req.Tenant, streams, zone, consumers)
answered, err := doRPCs(ctx, tenant, unanswered, zone, consumers)
if err != nil {
continue
}
responses = append(responses, resps...)
// Remove the answered streams from the slice. The slice of answered
// streams must be sorted so we can use sort.Search to subtract the
// two slices.
slices.Sort(answered)
streams = slices.DeleteFunc(streams, func(stream *proto.StreamMetadata) bool {
unanswered = slices.DeleteFunc(unanswered, func(stream *proto.StreamMetadata) bool {
// see https://pkg.go.dev/sort#Search
i := sort.Search(len(answered), func(i int) bool {
return answered[i] >= stream.StreamHash
})
return i < len(answered) && answered[i] == stream.StreamHash
})
}
// Any unanswered streams after exhausting all zones must be failed.
if len(streams) > 0 {
failed := make([]*proto.ExceedsLimitsResult, 0, len(streams))
for _, stream := range streams {
failed = append(failed, &proto.ExceedsLimitsResult{
StreamHash: stream.StreamHash,
Reason: uint32(limits.ReasonFailed),
})
}
responses = append(responses, &proto.ExceedsLimitsResponse{Results: failed})
}
return responses, nil
}

func (r *ringLimitsClient) doExceedsLimitsRPCs(ctx context.Context, tenant string, streams []*proto.StreamMetadata, zone string, consumers map[int32]string) ([]*proto.ExceedsLimitsResponse, []uint64, error) {
// For each stream, figure out which instance consume its partition.
instancesForStreams := make(map[string][]*proto.StreamMetadata)
for _, stream := range streams {
partition := int32(stream.StreamHash % uint64(r.numPartitions))
addr, ok := consumers[partition]
if !ok {
r.partitionsMissing.WithLabelValues(zone).Inc()
continue
}
instancesForStreams[addr] = append(instancesForStreams[addr], stream)
}
errg, ctx := errgroup.WithContext(ctx)
responseCh := make(chan *proto.ExceedsLimitsResponse, len(instancesForStreams))
answeredCh := make(chan uint64, len(streams))
for addr, streams := range instancesForStreams {
errg.Go(func() error {
client, err := r.pool.GetClientFor(addr)
if err != nil {
level.Error(r.logger).Log("msg", "failed to get client for instance", "instance", addr, "err", err.Error())
return nil
}
resp, err := client.(proto.IngestLimitsClient).ExceedsLimits(ctx, &proto.ExceedsLimitsRequest{
Tenant: tenant,
Streams: streams,
})
if err != nil {
level.Error(r.logger).Log("failed check execeed limits for instance", "instance", addr, "err", err.Error())
return nil
}
responseCh <- resp
for _, stream := range streams {
answeredCh <- stream.StreamHash
}
return nil
})
}
_ = errg.Wait()
close(responseCh)
close(answeredCh)
responses := make([]*proto.ExceedsLimitsResponse, 0, len(instancesForStreams))
for r := range responseCh {
responses = append(responses, r)
}
answered := make([]uint64, 0, len(streams))
for streamHash := range answeredCh {
answered = append(answered, streamHash)
}
return responses, answered, nil
return unanswered, nil
}

// allZones returns an iterator over all zones and the consumers for each
Expand All @@ -187,7 +290,7 @@ func (r *ringLimitsClient) doExceedsLimitsRPCs(ctx context.Context, tenant strin
// If ZoneAwarenessEnabled is false, it returns all partition consumers under
// a pseudo-zone ("").
func (r *ringLimitsClient) allZones(ctx context.Context) (iter.Seq2[string, map[int32]string], error) {
rs, err := r.ring.GetAllHealthy(LimitsRead)
rs, err := r.ring.GetAllHealthy(limitsRead)
if err != nil {
return nil, err
}
Expand Down Expand Up @@ -311,14 +414,29 @@ func (r *ringLimitsClient) getPartitionConsumers(ctx context.Context, instances
_ = errg.Wait()
close(responseCh)
highestTimestamp := make(map[int32]int64)
assigned := make(map[int32]string)
partitionConsumers := make(map[int32]string)
for resp := range responseCh {
for partition, assignedAt := range resp.response.AssignedPartitions {
if t := highestTimestamp[partition]; t < assignedAt {
highestTimestamp[partition] = assignedAt
assigned[partition] = resp.addr
partitionConsumers[partition] = resp.addr
}
}
}
return assigned, nil
return partitionConsumers, nil
}

func (r *ringLimitsClient) instancesForStreams(streams []*proto.StreamMetadata, zone string, instances map[int32]string) map[string][]*proto.StreamMetadata {
// For each stream, figure out which instance consume its partition.
instancesForStreams := make(map[string][]*proto.StreamMetadata)
for _, stream := range streams {
partition := int32(stream.StreamHash % uint64(r.numPartitions))
addr, ok := instances[partition]
if !ok {
r.partitionsMissing.WithLabelValues(zone).Inc()
continue
}
instancesForStreams[addr] = append(instancesForStreams[addr], stream)
}
return instancesForStreams
}
Loading