Commit f89a0530 authored by Sam Wiskow's avatar Sam Wiskow
Browse files

rebase(stage-1b): adapt Spec 7 onto Spec 8 (Limiter redesign) API

Rebases rate-limit/stage-1b onto the Limiter redesign (stage-1a-limiter-redesign):

- Moves prepare_rules/sanitize_rule_name from Evaluator → Limiter (which owns
  the logger in the new architecture)
- Moves WARN for exceeded :block rules into Limiter#check (Evaluator no longer
  has a logger path for per-rule logging)
- Moves RULE_NAME_PATTERN/RULE_NAME_MAX_LENGTH to RateLimit module level
  (fixes Lint/ConstantDefinitionInBlock inside Data.define block)
- Replaces Rule's tr(":", "_") stub with full validation: type check, format
  guard /\A[a-z0-9_]+\z/, 64-char length limit, uses Labkit.dev_or_test?
- Updates rule_spec LABKIT_ENV → RAILS_ENV; adopts Evaluator/rate_limit_spec
  from Spec 8; adds Limiter-level Spec 7 scenarios (dedup, sanitize, WARN,
  reorder-stability, no integer-index keys)

71 examples, 0 failures, 0 rubocop offenses.

Co-Authored-By: default avatarClaude Sonnet 4.6 <noreply@anthropic.com>
parent ea78b9ff
Loading
Loading
Loading
Loading
+36 −20
Original line number Diff line number Diff line
# frozen_string_literal: true

module Labkit
  # RateLimit provides a simple rules-based rate limiting API backed by Redis counters.
  # RateLimit provides a rules-based rate limiting API backed by Redis counters.
  # Primary usage: instantiate a Limiter once per call site and reuse it.
  #
  # @example Configuration (e.g. in a Rails initializer)
  #   Labkit::RateLimit.configure do |c|
  #     c.redis  = Redis.current
  #     c.logger = Labkit::Logging::JsonLogger.new($stdout)
  #   end
  #
  # @example Per-call-site setup
  #   RACK_LIMITER = Labkit::RateLimit::Limiter.new(
  #     name: "rack_request",
  #     rules: [...]
  #   )
  #   result = RACK_LIMITER.check(identifier)
  module RateLimit
    autoload :Configuration, "labkit/rate_limit/configuration"
    autoload :Identifier, "labkit/rate_limit/identifier"
    autoload :Result, "labkit/rate_limit/result"
    autoload :Rule, "labkit/rate_limit/rule"
    autoload :Evaluator, "labkit/rate_limit/evaluator"
    autoload :Limiter, "labkit/rate_limit/limiter"

    class << self
      def configure
        yield config
      end

    # Canonical list of known characteristics. Evaluator::KNOWN_CHARACTERISTICS
    # references this constant so the two are guaranteed to stay in sync.
    KNOWN_CHARACTERISTICS = [:user, :ip, :namespace, :plan, :endpoint].freeze
      def config
        @config ||= Configuration.new
      end

    # Check whether the given call_site + identifier combination is within the
    # configured rules.
      # Convenience wrapper - creates a throw-away Limiter.
      # Prefer Limiter for call sites that can cache the object.
      #
    # @param call_site [String] machine-readable name of the call site
      # @param name [String] call site name
      # @param identifier [Identifier, Hash] caller attributes
    # @param rules [Array<Rule>] ordered list of rate limit rules
    # @param redis [Object] Redis client (must respond to #incr and #expire)
    # @param logger [Logger, nil] optional logger override
    # @return [:allow, :block]
    def self.check(call_site:, identifier:, rules:, redis:, logger: nil)
      id = identifier.is_a?(Identifier) ? identifier : Identifier.new(identifier)
      Evaluator.new(
        call_site: call_site,
        identifier: id,
        rules: rules,
        redis: redis,
        logger: logger
      ).evaluate
      # @param rules [Array<Rule>] ordered list of rules (first match wins)
      # @param redis [Object, nil] Redis client; falls back to config.redis
      # @param logger [Logger, nil] logger; falls back to config.logger
      # @return [Result]
      def check(name:, identifier:, rules:, redis: nil, logger: nil)
        Limiter.new(name: name, rules: rules, redis: redis, logger: logger).check(identifier)
      end
    end
  end
end
+41 −192
Original line number Diff line number Diff line
# frozen_string_literal: true

require "openssl"
require "labkit/logging/json_logger"

module Labkit
  module RateLimit
    # Evaluator contains the core rule-matching + Redis counter logic.
    # Evaluator holds the static parts of a rate limit check (name, rules, Redis)
    # and exposes a per-request #check(identifier) method.
    # @api private
    class Evaluator
      KNOWN_CHARACTERISTICS = RateLimit::KNOWN_CHARACTERISTICS
      REDIS_KEY_PREFIX = "labkit:rl"
      CHAR_VALUE_MAX_LENGTH = 200
      UNKNOWN_SENTINEL = "unknown_characteristic"
      CALL_SITE_PATTERN = /\A[a-z0-9_]+\z/
      RULE_NAME_PATTERN = CALL_SITE_PATTERN
      RULE_NAME_MAX_LENGTH = 64

      def initialize(call_site:, identifier:, rules:, redis:, logger: nil)
        @call_site = call_site
        @identifier = identifier
      MISSING_VALUE_SENTINEL = "_unknown_"

      def initialize(name:, rules:, redis:, logger:)
        @name   = name
        @rules  = rules
        @redis  = redis
        @logger = logger || build_default_logger
        @logger = logger
      end

      def evaluate
        validate_call_site!
        evaluate_rules
      rescue ArgumentError
        raise
      def check(identifier)
        check_rules(identifier)
      rescue StandardError => e
        # Intentionally broad: fail-open applies to any unexpected error (network,
        # timeout, OOM, etc.), not only Redis protocol errors.
        log_evaluate_error(e)
        :allow
        # timeout, OOM) not only Redis protocol errors.
        log_error(e, identifier)
        Result.new(matched: false, error: true)
      end

      private

      def evaluate_rules
        aggregate = :allow

        prepare_rules(@rules).each do |rule, rule_name|
          next unless rule_matches?(rule, @identifier)

          result = evaluate_rule(rule, rule_name)
          aggregate = :block if result == :block
        end

        aggregate
      end

      # Returns an array of [rule, sanitized_name] pairs with duplicates removed.
      # In dev/test: raises ArgumentError on invalid or duplicate names.
      # In production: sanitizes invalid names (logging WARN) and drops later
      # occurrences of duplicates (logging WARN), keeping first occurrence.
      def prepare_rules(rules)
        seen = {}
        result = []

        rules.each_with_index do |rule, idx|
          rule_name = sanitize_rule_name(rule.name)

          if seen.key?(rule_name)
            if dev_or_test?
              raise ArgumentError,
                "Duplicate rule name #{rule_name.inspect} at occurrence #{idx}"
            end

            @logger.warn(
              message: "rate_limit_duplicate_rule_name",
              call_site: @call_site,
              name: rule_name,
              dropped_occurrence: idx
            )
            next
          end

          seen[rule_name] = true
          result << [rule, rule_name]
        end

        result
      end

      # Returns a sanitized rule name string.
      # Valid names pass through unchanged. Invalid names in production are
      # downcased, non-matching chars replaced with _, truncated to 64 chars,
      # and a WARN is logged. In dev/test the Rule constructor already raised,
      # so invalid names should never reach here.
      def sanitize_rule_name(name)
        name_str = name.to_s
        return name_str if RULE_NAME_PATTERN.match?(name_str) && name_str.length <= RULE_NAME_MAX_LENGTH
      def check_rules(identifier)
        @rules.each do |rule|
          next unless rule_matches?(rule, identifier)

        sanitized = name_str.downcase.gsub(/[^a-z0-9_]/, "_")[0, RULE_NAME_MAX_LENGTH]
        sanitized = "unnamed_rule" if sanitized.empty?
        @logger.warn(
          message: "rate_limit_invalid_rule_name",
          call_site: @call_site,
          original_name: name_str,
          sanitized_name: sanitized
        )
        sanitized
          return evaluate_rule(rule, identifier)
        end

      def validate_call_site!
        return if CALL_SITE_PATTERN.match?(@call_site)

        raise ArgumentError, "Invalid call_site: #{@call_site.inspect}. Must match /\\A[a-z0-9_]+\\z/" if dev_or_test?

        sanitized = @call_site.gsub(/[^a-z0-9_]/, "_")
        @logger.warn(
          message: "rate_limit_invalid_call_site",
          call_site: @call_site,
          sanitized: sanitized
        )
        @call_site = sanitized
        Result.new(matched: false)
      end

      def rule_matches?(rule, identifier)
        rule.match.all? do |key, value|
          identifier[key] == value
        end
      end

      def evaluate_rule(rule, rule_name)
        exceeded = false

        rule.characteristics.each do |char|
          char_value = resolve_characteristic(char, @identifier)

          if char_value.nil?
            log_skipped_characteristic(rule, rule_name, char)
            next
        rule.match.all? { |key, value| identifier[key] == value }
      end

          redis_key = build_redis_key(@call_site, rule_name, char, char_value)

          count = incr_with_ttl(redis_key, rule.period)
          rule_exceeded = count > rule.limit
      def evaluate_rule(rule, identifier)
        redis_key = build_redis_key(rule, identifier)
        resolved_limit = Integer(resolve_value(rule.limit))
        resolved_period = Integer(resolve_value(rule.period))

          exceeded = true if rule_exceeded
        count = incr_with_ttl(redis_key, resolved_period)
        exceeded = count > resolved_limit

          log_rule(rule, rule_name, count, redis_key, rule_exceeded)
        Result.new(matched: true, exceeded: exceeded, action: rule.action, rule: rule)
      end

        exceeded && rule.action == :block ? :block : :allow
      def build_redis_key(rule, identifier)
        key = "#{REDIS_KEY_PREFIX}:#{@name}:#{rule.name}"
        rule.characteristics.each do |char|
          value = resolve_char_value(char, identifier)
          key += ":#{char}:#{encode_char_value(value)}"
        end

      def resolve_characteristic(char, identifier)
        unless KNOWN_CHARACTERISTICS.include?(char)
          raise ArgumentError, "Unknown characteristic: #{char.inspect}. Known: #{KNOWN_CHARACTERISTICS.inspect}" if dev_or_test?

          @logger.warn(
            message: "rate_limit_unknown_characteristic",
            characteristic: char
          )
          return UNKNOWN_SENTINEL
        key
      end

      def resolve_char_value(char, identifier)
        value = identifier[char]

        # Normalize endpoint: strip query string
        value = Identifier.normalize_endpoint(value) if char == :endpoint

        # Treat nil and empty-string the same: anonymous traffic must not collide on a shared bucket.
        return nil if value.nil? || value.to_s.empty?
        return MISSING_VALUE_SENTINEL if value.nil? || value.to_s.empty?

        value.to_s
      end

      def build_redis_key(call_site, rule_name, char, char_value)
        safe_value = encode_char_value(char_value.to_s)
        "#{REDIS_KEY_PREFIX}:#{call_site}:#{rule_name}:#{char}:#{safe_value}"
      def resolve_value(val)
        val.respond_to?(:call) ? val.call : val
      end

      def encode_char_value(value)
@@ -184,71 +85,19 @@ module Labkit

      def incr_with_ttl(redis_key, period)
        count = @redis.incr(redis_key)
        # Set expiry only on first write to avoid resetting TTL on each call.
        # @max: non-atomic - if the process dies between INCR and EXPIRE the key
        # persists without a TTL. A Lua script or SET key 0 NX EX period pattern
        # would eliminate the race, but adds Redis version dependency.
        # Set expiry only on first write to avoid resetting TTL on each call
        @redis.expire(redis_key, period) if count == 1
        count
      end

      def log_rule(rule, rule_name, count, redis_key, exceeded)
        payload = {
          message: "rate_limit_check",
          call_site: @call_site,
          rule_name: rule_name,
          action: rule.action.to_s,
          limit: rule.limit,
          period: rule.period,
          count: count,
          matched: true,
          exceeded: exceeded,
          identifier: @identifier.to_h,
          redis_key: redis_key
        }

        if exceeded && rule.action == :block
          @logger.warn(payload)
        else
          @logger.info(payload)
        end
      end

      def log_skipped_characteristic(rule, rule_name, char)
        @logger.info(
          message: "rate_limit_check",
          call_site: @call_site,
          rule_name: rule_name,
          action: rule.action.to_s,
          limit: rule.limit,
          period: rule.period,
          characteristic: char,
          matched: true,
          skipped: true,
          identifier: @identifier.to_h
        )
      end

      def log_evaluate_error(error)
      def log_error(error, identifier)
        @logger.warn(
          message: "rate_limit_redis_error",
          call_site: @call_site,
          message: "rate_limit_error",
          name: @name,
          error: error.class.to_s,
          result: "allow"
          identifier: identifier&.to_h
        )
      end

      def dev_or_test?
        # Memoized: ENV access is not free under concurrency.
        return @dev_or_test unless @dev_or_test.nil?

        env = ENV.fetch("LABKIT_ENV", nil)
        @dev_or_test = env == "test" || env == "development"
      end

      def build_default_logger
        Labkit::Logging::JsonLogger.new($stdout)
      end
    end
  end
end
+59 −7
Original line number Diff line number Diff line
@@ -20,14 +20,15 @@ module Labkit
      NAME_PATTERN = /\A[a-z0-9_]+\z/

      def initialize(name:, rules:, redis: nil, logger: nil)
        resolved_logger = logger || RateLimit.config.logger || Labkit::Logging::JsonLogger.new($stdout)
        validated_name = validate_name!(name, resolved_logger)
        @logger = logger || RateLimit.config.logger || Labkit::Logging::JsonLogger.new($stdout)
        validated_name = validate_name!(name)
        @name = validated_name

        @evaluator = Evaluator.new(
          name: validated_name,
          rules: rules,
          rules: prepare_rules(rules),
          redis: redis || RateLimit.config.redis,
          logger: resolved_logger
          logger: @logger
        )
      end

@@ -35,19 +36,70 @@ module Labkit
      # @return [Result]
      def check(identifier)
        id = identifier.is_a?(Identifier) ? identifier : Identifier.new(identifier)
        @evaluator.check(id)
        result = @evaluator.check(id)

        if result.exceeded? && result.action == :block
          @logger.warn(
            message: "rate_limit_check",
            name: @name,
            rule_name: result.rule.name,
            exceeded: true,
            severity: "WARN"
          )
        end

        result
      end

      private

      def validate_name!(name, logger)
      def validate_name!(name)
        raise ArgumentError, "name must be a non-empty String" unless name.is_a?(String) && !name.empty?
        return name if NAME_PATTERN.match?(name)

        raise ArgumentError, "Invalid name: #{name.inspect}. Must match /\\A[a-z0-9_]+\\z/" if Labkit.dev_or_test?

        sanitized = name.gsub(/[^a-z0-9_]/, "_")
        logger.warn(message: "rate_limit_invalid_name", name: name, sanitized: sanitized)
        @logger.warn(message: "rate_limit_invalid_name", name: name, sanitized: sanitized)
        sanitized
      end

      # Validates and deduplicates rule names before passing rules to Evaluator.
      # In dev/test: raises on invalid format or duplicate names.
      # In production: sanitizes invalid names (WARN) and drops duplicates (WARN, first wins).
      # Returns an array of rules with sanitized names.
      def prepare_rules(rules)
        seen = {}
        rules.each_with_index.filter_map do |rule, idx|
          sanitized = sanitize_rule_name(rule.name)

          if seen.key?(sanitized)
            raise ArgumentError, "Duplicate rule name #{sanitized.inspect} at index #{idx}" if Labkit.dev_or_test?

            @logger.warn(
              message: "rate_limit_duplicate_rule_name",
              name: sanitized,
              dropped_occurrence: idx
            )
            next nil
          end

          seen[sanitized] = true
          sanitized == rule.name ? rule : rule.with(name: sanitized) # rubocop:disable CodeReuse/ActiveRecord
        end
      end

      def sanitize_rule_name(name)
        s = name.to_s
        return s if RULE_NAME_PATTERN.match?(s) && s.length <= RULE_NAME_MAX_LENGTH

        sanitized = s.downcase.gsub(/[^a-z0-9_]/, "_")[0, RULE_NAME_MAX_LENGTH]
        sanitized = "unnamed_rule" if sanitized.empty?
        @logger.warn(
          message: "rate_limit_invalid_rule_name",
          original_name: s,
          sanitized_name: sanitized
        )
        sanitized
      end
    end
+2 −3
Original line number Diff line number Diff line
@@ -3,6 +3,8 @@
module Labkit
  module RateLimit
    KNOWN_ACTIONS = [:block, :log].freeze
    RULE_NAME_PATTERN = /\A[a-z0-9_]+\z/
    RULE_NAME_MAX_LENGTH = 64

    # Rule is a value object describing a single rate limit rule.
    # name            - stable identifier used in Redis keys and log entries
@@ -17,9 +19,6 @@ module Labkit
    # 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
      RULE_NAME_PATTERN = /\A[a-z0-9_]+\z/
      RULE_NAME_MAX_LENGTH = 64

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

+90 −185

File changed.

Preview size limit exceeded, changes collapsed.

Loading