Verified Commit b283c27c authored by Bob Van Landuyt's avatar Bob Van Landuyt 💬 Committed by GitLab
Browse files

Merge branch 'rate-limit/unique-cardinality' into 'master'

feat(rate_limit): SET-mode check_unique / peek_unique primitive

See merge request !292

Merged-by: Bob Van Landuyt's avatarBob Van Landuyt <bob@gitlab.com>
Approved-by: Bob Van Landuyt's avatarBob Van Landuyt <bob@gitlab.com>
Reviewed-by: Bob Van Landuyt's avatarBob Van Landuyt <bob@gitlab.com>
Reviewed-by: Max Woolf's avatarMax Woolf <mwoolf@gitlab.com>
Reviewed-by: default avatarGitLab Duo <gitlab-duo@gitlab.com>
Co-authored-by: Max Woolf's avatarMax Woolf <mwoolf@gitlab.com>
parents 663989ec ef148a80
Loading
Loading
Loading
Loading
Loading
+85 −2
Original line number Diff line number Diff line
@@ -38,6 +38,26 @@ module Labkit
        return {count, redis.call('TTL', KEYS[1])}
      LUA

      # Atomic SADD + SCARD + conditional EXPIRE. SET-cardinality counterpart
      # of INCR_SCRIPT; same shape (read TTL, mutate, set TTL when missing,
      # return post-state {count, TTL}). count is SCARD, not the SADD return.
      #
      # ttl_before < 0 covers TTL=-2 (key missing) and TTL=-1 (no expiry),
      # so this also self-heals orphan keys left without TTL.
      SADD_SCRIPT = Labkit::Redis::Script.new(<<~LUA)
        local ttl = ARGV[1]
        local member = ARGV[2]
        local ttl_before = redis.call('TTL', KEYS[1])

        redis.call('SADD', KEYS[1], member)
        local count = redis.call('SCARD', KEYS[1])
        if ttl_before < 0 then
          redis.call('EXPIRE', KEYS[1], ttl)
        end

        return {count, redis.call('TTL', KEYS[1])}
      LUA

      def initialize(name:, rules:, redis:, logger:)
        @name   = name
        @rules  = rules
@@ -70,10 +90,21 @@ module Labkit

      # :log rules are non-terminating: they emit metrics and continue,
      # so a shadow :log rule cannot disable a following :block rule.
      #
      # SET-mode rules (rule.count_distinct set) that match but whose identifier
      # is missing the count_distinct key fail open + log + bump errors_total, and
      # the loop continues to the next rule (the rule is treated as not applicable
      # rather than aborting the whole evaluation).
      def check_rules(identifier, cost)
        @rules.each do |rule|
          next unless rule_matches?(rule, identifier)

          if rule.count_distinct && missing_count_distinct_value?(rule, identifier)
            log_missing_count_distinct(rule, identifier)
            report_error_metrics
            next
          end

          result = evaluate_rule(rule, identifier, cost)
          report_matched_metrics(result)
          return result unless rule.action == :log
@@ -85,6 +116,10 @@ module Labkit

      # Mirror of check_rules without metrics: peek skips :log rules (their state
      # is unobservable through peek).
      #
      # peek does not need the count_distinct identifier key - it reads SCARD on
      # the rule-keyed compound key, which contains the cardinality across all
      # members. So missing-key fail-open does not apply here.
      def peek_rules(identifier)
        @rules.each do |rule|
          next if rule.action == :log
@@ -100,12 +135,25 @@ module Labkit
        rule.match.all? { |key, matcher| matcher.match?(identifier[key]) }
      end

      def missing_count_distinct_value?(rule, identifier)
        value = identifier[rule.count_distinct]
        value.nil? || value.to_s.empty?
      end

      def evaluate_rule(rule, identifier, cost)
        redis_key = build_redis_key(rule, identifier)
        resolved_limit = Integer(resolve_value(rule.limit))
        resolved_period = Integer(resolve_value(rule.period))

        count, ttl = incr_with_ttl(redis_key, resolved_period, cost)
        # cost is ignored for count_distinct rules: SADD is binary (a member is
        # either added or not), and the post-add count is SCARD regardless.
        count, ttl =
          if rule.count_distinct
            sadd_with_ttl(redis_key, identifier[rule.count_distinct], resolved_period)
          else
            incr_with_ttl(redis_key, resolved_period, cost)
          end

        build_result(rule, resolved_limit, resolved_period, count, ttl)
      end

@@ -114,7 +162,7 @@ module Labkit
        resolved_limit = Integer(resolve_value(rule.limit))
        resolved_period = Integer(resolve_value(rule.period))

        count, ttl = read_with_ttl(redis_key)
        count, ttl = rule.count_distinct ? scard_with_ttl(redis_key) : read_with_ttl(redis_key)
        build_result(rule, resolved_limit, resolved_period, count, ttl)
      end

@@ -191,6 +239,31 @@ module Labkit
        end
      end

      # Atomic SADD + SCARD + conditional EXPIRE in one Redis operation via Lua.
      # See SADD_SCRIPT for the body. Mirrors incr_with_ttl's shape, including
      # the Float-typed count for uniformity with the INCR path.
      def sadd_with_ttl(redis_key, member, period)
        member_str = encode_char_value(member.to_s)
        @redis.with do |conn|
          raw_count, ttl = SADD_SCRIPT.eval(conn, keys: [redis_key], argv: [period, member_str])
          [Float(raw_count), ttl]
        end
      end

      # Pipelined SCARD + TTL. SCARD on a missing key returns 0, so no
      # explicit nil handling is needed (unlike GET in read_with_ttl).
      # SCARD is integer-valued; coerced to Float for type uniformity with
      # the INCR path so callers see a consistent count type.
      def scard_with_ttl(redis_key)
        @redis.with do |conn|
          scard, ttl = conn.pipelined do |pipe|
            pipe.scard(redis_key)
            pipe.ttl(redis_key)
          end
          [Float(scard), ttl]
        end
      end

      def log_error(error, identifier)
        @logger.warn(
          message: "rate_limit_error",
@@ -200,6 +273,16 @@ module Labkit
        )
      end

      def log_missing_count_distinct(rule, identifier)
        @logger.warn(
          message: "rate_limit_missing_count_distinct",
          name: @name,
          rule: rule.name,
          count_distinct: rule.count_distinct.to_s,
          identifier: identifier&.to_h
        )
      end

      def report_matched_metrics(result)
        Metrics.calls_total.increment(
          rate_limiter: @name,
+26 −3
Original line number Diff line number Diff line
@@ -17,12 +17,32 @@ module Labkit
    #                   (count but always permit; terminates evaluation on match
    #                   regardless of whether the limit was exceeded)
    # characteristics - identifier keys used to build the compound Redis counter key
    # count_distinct  - optional Symbol naming an identifier key. When set, the rule
    #                   counts the number of distinct values seen for that key within
    #                   the (characteristics-bucketed) period, backed by a Redis SET.
    #                   When nil (default), the rule counts the number of calls,
    #                   backed by INCR. The named key must not overlap +characteristics+.
    #
    # +name+ must be a lowercase alphanumeric-and-underscore string of at most 64
    # characters. It is used as the middle segment of every Redis counter key for
    # this rule, so changing a rule's name mid-window abandons its in-flight counters.
    Rule = Data.define(:name, :match, :limit, :period, :action, :characteristics) do
      def initialize(name:, limit:, period:, characteristics:, match: {}, action: :block)
    Rule = Data.define(:name, :match, :limit, :period, :action, :characteristics, :count_distinct) do
      def self.normalize_count_distinct(value, characteristics_arr)
        sym =
          case value
          when nil    then nil
          when Symbol then value
          when String then value.to_sym
          else
            raise ArgumentError, "count_distinct must be a Symbol, String, or nil, got #{value.class}"
          end

        raise ArgumentError, "count_distinct #{sym.inspect} must not overlap characteristics #{characteristics_arr.inspect}" if sym && characteristics_arr.include?(sym)

        sym
      end

      def initialize(name:, limit:, period:, characteristics:, match: {}, action: :block, count_distinct: nil)
        raise ArgumentError, "name must be a String or Symbol, got #{name.class}" unless name.is_a?(String) || name.is_a?(Symbol)

        name_str = name.to_s
@@ -36,13 +56,16 @@ module Labkit
          raise ArgumentError, "Rule name too long: #{name.inspect}. Maximum 64 characters" if name_str.length > RULE_NAME_MAX_LENGTH
        end

        characteristics_arr = Array(characteristics).map(&:to_sym).freeze

        super(
          name: name_str.freeze,
          match: match.transform_keys(&:to_sym).transform_values { |v| Matcher.build(v) }.freeze,
          limit: limit,
          period: period,
          action: action_sym,
          characteristics: Array(characteristics).map(&:to_sym).freeze
          characteristics: characteristics_arr,
          count_distinct: self.class.normalize_count_distinct(count_distinct, characteristics_arr)
        )
      end
    end
+272 −2
Original line number Diff line number Diff line
@@ -12,10 +12,10 @@ RSpec.describe Labkit::RateLimit::Evaluator do
  let(:null_logger) { instance_double(Labkit::Logging::JsonLogger, warn: nil, error: nil) }
  let(:identifier) { Labkit::RateLimit::Identifier.new(user: 42, ip: "1.2.3.4") }

  def make_rule(name: "default", match: {}, limit: 100, period: 60, action: :block, characteristics: [:user])
  def make_rule(name: "default", match: {}, limit: 100, period: 60, action: :block, characteristics: [:user], count_distinct: nil)
    Labkit::RateLimit::Rule.new(
      name: name, match: match, limit: limit, period: period,
      action: action, characteristics: characteristics
      action: action, characteristics: characteristics, count_distinct: count_distinct
    )
  end

@@ -813,4 +813,274 @@ RSpec.describe Labkit::RateLimit::Evaluator do
      expect(raw_redis.keys("labkit:rl:rack_request:block_rule_l:*")).to be_empty
    end
  end

  describe "#check with count_distinct (SET-mode)" do
    let(:unique_id) { Labkit::RateLimit::Identifier.new(user: 42, ip: "1.2.3.4", project: 99) }

    it "SADDs the count_distinct value to the rule-keyed compound key" do
      rule = make_rule(name: "uniq_rule", characteristics: [:user], count_distinct: :project)
      key = "labkit:rl:rack_request:uniq_rule:user:42"

      evaluator(rules: [rule]).check(unique_id)

      expect(raw_redis.sismember(key, "99")).to be(true)
      expect(raw_redis.scard(key)).to eq(1)
    end

    it "coerces non-string count_distinct values via to_s" do
      rule = make_rule(name: "uniq_rule", characteristics: [:user], count_distinct: :project)
      key = "labkit:rl:rack_request:uniq_rule:user:42"

      evaluator(rules: [rule]).check(unique_id) # project: 99 (Integer)

      expect(raw_redis.sismember(key, "99")).to be(true)
    end

    it "SHA-256-encodes count_distinct values longer than CHAR_VALUE_MAX_LENGTH" do
      long_value = "x" * 201
      expected_hash = OpenSSL::Digest::SHA256.hexdigest(long_value)
      rule = make_rule(name: "uniq_rule", characteristics: [:user], count_distinct: :project)
      id = Labkit::RateLimit::Identifier.new(user: 42, project: long_value)
      key = "labkit:rl:rack_request:uniq_rule:user:42"

      evaluator(rules: [rule]).check(id)

      expect(raw_redis.sismember(key, expected_hash)).to be(true)
      expect(raw_redis.sismember(key, long_value)).to be(false)
    end

    it "sets the TTL on first write" do
      rule = make_rule(name: "uniq_rule", period: 120, count_distinct: :project)
      key = "labkit:rl:rack_request:uniq_rule:user:42"

      evaluator(rules: [rule]).check(unique_id)

      expect(raw_redis.ttl(key)).to be_between(1, 120)
    end

    it "does not reset TTL on subsequent writes to a key that already has one" do
      rule = make_rule(name: "uniq_rule", period: 120, count_distinct: :project)
      key = "labkit:rl:rack_request:uniq_rule:user:42"
      ev = evaluator(rules: [rule])

      ev.check(unique_id)
      # Force a distinctive TTL the rule's period would never produce; a subsequent
      # check that mistakenly EXPIRE'd would clobber it back to ~120.
      raw_redis.expire(key, 7)

      ev.check(Labkit::RateLimit::Identifier.new(user: 42, project: 100))

      expect(raw_redis.ttl(key)).to be_between(1, 7)
      expect(raw_redis.scard(key)).to eq(2)
    end

    it "self-heals a key that exists without expiry (TTL = -1)" do
      rule = make_rule(name: "heal_rule", period: 60, count_distinct: :project)
      key = "labkit:rl:rack_request:heal_rule:user:42"
      raw_redis.sadd(key, %w[1 2]) # no TTL; simulates an orphan from a prior bug

      evaluator(rules: [rule]).check(unique_id)

      expect(raw_redis.ttl(key)).to be_between(1, 60)
      expect(raw_redis.scard(key)).to eq(3)
    end

    it "reports exceeded when the post-add cardinality exceeds the limit" do
      rule = make_rule(name: "uniq_rule", limit: 5, action: :block, count_distinct: :project)
      key = "labkit:rl:rack_request:uniq_rule:user:42"
      raw_redis.sadd(key, %w[a b c d e]) # already at the limit; project 99 tips us over

      result = evaluator(rules: [rule]).check(unique_id)

      expect(result.exceeded?).to be(true)
      expect(result.action).to eq(:block)
      expect(result.info.count).to eq(6.0)
    end

    it "fails open on Redis error", :aggregate_failures do
      broken = instance_double(Redis)
      allow(broken).to receive(:evalsha).and_raise(RuntimeError, "connection refused")
      broken_pool = PooledRedis.new(broken)
      logger = instance_double(Labkit::Logging::JsonLogger)
      expect(logger).to receive(:warn).with(
        hash_including(message: "rate_limit_error", error: "RuntimeError")
      )

      rule = make_rule(name: "uniq_rule", count_distinct: :project)
      result = described_class.new(name: "rack_request", rules: [rule], redis: broken_pool, logger: logger)
        .check(unique_id)

      expect(result.error?).to be(true)
      expect(result.exceeded?).to be(false)
    end

    it "still reports exceeded on re-add when the set is already over the limit" do
      rule = make_rule(name: "uniq_rule", limit: 5, action: :block, count_distinct: :project)
      key = "labkit:rl:rack_request:uniq_rule:user:42"
      raw_redis.sadd(key, %w[a b c d e 99]) # 6 members; 99 already present

      result = evaluator(rules: [rule]).check(unique_id) # re-adds 99, scard unchanged

      expect(result.exceeded?).to be(true)
      expect(result.action).to eq(:block)
      expect(result.info.count).to eq(6.0)
    end

    it "treats a :log rule as non-terminating and continues to the next rule" do
      log_r = make_rule(name: "log_r", action: :log, limit: 1, count_distinct: :project)
      block_r = make_rule(name: "block_r", action: :block, limit: 100, count_distinct: :project)

      result = evaluator(rules: [log_r, block_r]).check(unique_id)

      expect(result.rule).to eq(block_r)
      expect(result.action).to eq(:allow)
      expect(raw_redis.scard("labkit:rl:rack_request:log_r:user:42")).to eq(1)
      expect(raw_redis.scard("labkit:rl:rack_request:block_r:user:42")).to eq(1)
    end

    it "uses _unknown_ sentinel for missing characteristics in the rule key (count_distinct is separate)" do
      rule = make_rule(name: "uniq_rule", characteristics: [:user, :ip], count_distinct: :project)
      id = Labkit::RateLimit::Identifier.new(user: 42, project: 99)
      key = "labkit:rl:rack_request:uniq_rule:user:42:ip:_unknown_"

      evaluator(rules: [rule]).check(id)

      expect(raw_redis.sismember(key, "99")).to be(true)
    end

    describe "fail-open + log when count_distinct identifier key is missing" do
      it "logs rate_limit_missing_count_distinct and continues to the next rule when the key is absent" do
        rule = make_rule(name: "uniq_rule", count_distinct: :project)
        id = Labkit::RateLimit::Identifier.new(user: 42)
        logger = instance_double(Labkit::Logging::JsonLogger)
        expect(logger).to receive(:warn).with(
          hash_including(
            message: "rate_limit_missing_count_distinct",
            rule: "uniq_rule",
            count_distinct: "project"
          )
        )

        result = described_class.new(name: "rack_request", rules: [rule], redis: redis, logger: logger).check(id)

        expect(result.matched?).to be(false)
        expect(result.action).to eq(:allow)
      end

      it "logs when the value is nil" do
        rule = make_rule(name: "uniq_rule", count_distinct: :project)
        id = Labkit::RateLimit::Identifier.new(user: 42, project: nil)

        expect(null_logger).to receive(:warn).with(hash_including(message: "rate_limit_missing_count_distinct"))
        evaluator(rules: [rule]).check(id)
      end

      it "logs when the value is an empty string" do
        rule = make_rule(name: "uniq_rule", count_distinct: :project)
        id = Labkit::RateLimit::Identifier.new(user: 42, project: "")

        expect(null_logger).to receive(:warn).with(hash_including(message: "rate_limit_missing_count_distinct"))
        evaluator(rules: [rule]).check(id)
      end

      it "does not create the Redis key when the count_distinct value is missing" do
        rule = make_rule(name: "uniq_rule", count_distinct: :project)
        id = Labkit::RateLimit::Identifier.new(user: 42)
        key = "labkit:rl:rack_request:uniq_rule:user:42"

        evaluator(rules: [rule]).check(id)

        expect(raw_redis.exists?(key)).to be(false)
      end

      it "falls through to a following matching rule when the SET-mode rule's key is missing" do
        # SET-mode rule comes first; its key is missing, so it's skipped.
        # A subsequent INCR-mode rule then applies normally against real Redis.
        uniq = make_rule(name: "uniq_first", count_distinct: :project, limit: 5)
        plain = make_rule(name: "plain_after", limit: 10, action: :block)
        id = Labkit::RateLimit::Identifier.new(user: 42)

        result = evaluator(rules: [uniq, plain]).check(id)

        expect(result.rule).to eq(plain)
        expect(result.matched?).to be(true)
        expect(stored_count("labkit:rl:rack_request:plain_after:user:42")).to eq(1.0)
      end

      it "bumps errors_total when a count_distinct rule is skipped" do
        rule = make_rule(name: "uniq_rule", count_distinct: :project)
        id = Labkit::RateLimit::Identifier.new(user: 42)

        expect(Labkit::RateLimit::Metrics.errors_total)
          .to receive(:increment).with(rate_limiter: "rack_request")
        allow(Labkit::RateLimit::Metrics.calls_total).to receive(:increment)

        evaluator(rules: [rule]).check(id)
      end
    end
  end

  describe "#peek with count_distinct (SET-mode)" do
    it "reads SCARD without mutating the set or touching TTL" do
      rule = make_rule(name: "uniq_rule", characteristics: [:user], count_distinct: :project)
      key = "labkit:rl:rack_request:uniq_rule:user:42"
      raw_redis.sadd(key, %w[a b c])
      raw_redis.expire(key, 7)

      result = evaluator(rules: [rule]).peek(identifier)

      expect(result.matched?).to be(true)
      expect(result.info.count).to eq(3.0)
      expect(raw_redis.scard(key)).to eq(3) # no SADD happened
      expect(raw_redis.ttl(key)).to be_between(1, 7) # TTL not extended
    end

    it "reports count=0 when key is missing (SCARD returns 0)" do
      rule = make_rule(name: "uniq_rule", limit: 5, period: 60, count_distinct: :project)

      result = evaluator(rules: [rule]).peek(identifier)

      expect(result.matched?).to be(true)
      expect(result.exceeded?).to be(false)
      expect(result.info.count).to eq(0.0)
      expect(result.info.remaining).to eq(5)
    end

    it "reports exceeded when the current cardinality is over the limit" do
      rule = make_rule(name: "uniq_rule", limit: 5, action: :block, count_distinct: :project)
      key = "labkit:rl:rack_request:uniq_rule:user:42"
      raw_redis.sadd(key, %w[a b c d e f g h i j])

      result = evaluator(rules: [rule]).peek(identifier)

      expect(result.exceeded?).to be(true)
      expect(result.action).to eq(:block)
    end

    it "fails open on Redis error" do
      broken = instance_double(Redis)
      allow(broken).to receive(:pipelined).and_raise(RuntimeError, "down")
      broken_pool = PooledRedis.new(broken)
      rule = make_rule(name: "uniq_rule", count_distinct: :project)

      result = described_class.new(name: "rack_request", rules: [rule], redis: broken_pool, logger: null_logger)
        .peek(identifier)

      expect(result.error?).to be(true)
      expect(result.exceeded?).to be(false)
    end

    it "does not require the count_distinct identifier key (reads the bucket cardinality)" do
      rule = make_rule(name: "uniq_rule", characteristics: [:user], count_distinct: :project)
      id = Labkit::RateLimit::Identifier.new(user: 42) # no :project
      key = "labkit:rl:rack_request:uniq_rule:user:42"
      raw_redis.sadd(key, %w[a b c])

      expect(null_logger).not_to receive(:warn)
      result = evaluator(rules: [rule]).peek(id)

      expect(result.matched?).to be(true)
      expect(result.info.count).to eq(3.0)
    end
  end
end
+39 −0
Original line number Diff line number Diff line
@@ -235,4 +235,43 @@ RSpec.describe Labkit::RateLimit::Limiter do
      end
    end
  end

  describe "#check with a count_distinct rule" do
    before do
      # SADD_SCRIPT.eval returns [scard, ttl]
      allow(raw_redis).to receive(:evalsha).and_return([4, 30])
    end

    it "returns a Result whose count is the SCARD post-add" do
      r = Labkit::RateLimit::Rule.new(
        name: "uniq", limit: 10, period: 60,
        characteristics: [:user], count_distinct: :project
      )

      result = described_class.new(name: "rack_request", rules: [r], redis: redis, logger: logger)
        .check({ user: 42, project: 99 })

      expect(result.matched?).to be(true)
      expect(result.info.count).to eq(4.0)
      expect(result.info.remaining).to eq(6)
    end
  end

  describe "#peek with a count_distinct rule" do
    before do
      allow(raw_redis).to receive(:pipelined).and_return([0, -2])
    end

    it "does not invoke the SADD script" do
      r = Labkit::RateLimit::Rule.new(
        name: "uniq", limit: 10, period: 60,
        characteristics: [:user], count_distinct: :project
      )
      lim = described_class.new(name: "rack_request", rules: [r], redis: redis, logger: logger)

      expect(raw_redis).not_to receive(:evalsha)
      expect(raw_redis).not_to receive(:expire)
      lim.peek({ user: 42 })
    end
  end
end
+32 −0
Original line number Diff line number Diff line
@@ -235,4 +235,36 @@ RSpec.describe Labkit::RateLimit::Rule do
      expect(valid_rule.characteristics).to be_frozen
    end
  end

  describe "count_distinct" do
    it "defaults to nil" do
      expect(valid_rule.count_distinct).to be_nil
    end

    it "accepts a Symbol" do
      rule = valid_rule(count_distinct: :project_id)
      expect(rule.count_distinct).to eq(:project_id)
    end

    it "coerces a String to a Symbol" do
      rule = valid_rule(count_distinct: "project_id")
      expect(rule.count_distinct).to eq(:project_id)
    end

    it "raises when set to a non-Symbol, non-String, non-nil value" do
      expect { valid_rule(count_distinct: 42) }
        .to raise_error(ArgumentError, /count_distinct must be a Symbol, String, or nil/)
    end

    it "raises when it overlaps characteristics" do
      expect { valid_rule(characteristics: [:user, :project_id], count_distinct: :project_id) }
        .to raise_error(ArgumentError, /must not overlap characteristics/)
    end

    it "allows a value that does not overlap characteristics" do
      rule = valid_rule(characteristics: [:user, :namespace_id], count_distinct: :project_id)
      expect(rule.count_distinct).to eq(:project_id)
      expect(rule.characteristics).to eq([:user, :namespace_id])
    end
  end
end
Loading