Commit 6d96a062 authored by Sam Wiskow's avatar Sam Wiskow
Browse files

fix: resolve rubocop violations in rate_limit implementation

- rescue StandardError instead of bare rescue (Style/RescueStandardError)
- Refactor evaluate to avoid return from block; fail-open rescue moved
  to outer method (Cop/AvoidReturnFromBlocks)
- Use guard clauses in validate_call_site! and resolve_characteristic
  (Style/GuardClause)
- Rename increment_counter -> incr_with_ttl to avoid Rails/SkipsModelValidations
  false positive
- Remove parentheses from ternary condition (Style/TernaryParentheses)
- Replace Digest::SHA256 with OpenSSL::Digest::SHA256 for FIPS compliance
  (Fips/OpenSSL)
- Use ENV.fetch instead of ENV[] (Style/FetchEnvVar)
- Move normalize_endpoint class method before instance methods in Identifier
  (Layout/ClassStructure)
parent d23c7495
Loading
Loading
Loading
Loading
+21 −24
Original line number Diff line number Diff line
# frozen_string_literal: true

require "digest"
require "json"
require "logger"
require "openssl"

module Labkit
  module RateLimit
@@ -24,34 +24,34 @@ module Labkit

      def evaluate
        validate_call_site!
        evaluate_rules
      rescue ArgumentError
        raise
      rescue StandardError => e
        log_redis_error(e)
        :allow
      end

      private

      def evaluate_rules
        aggregate = :allow

        @rules.each_with_index do |rule, index|
          next unless rule_matches?(rule, @identifier)

          begin
          result = evaluate_rule(rule, index)
          aggregate = :block if result == :block
          rescue ArgumentError
            raise
          rescue => e
            log_redis_error(e)
            return :allow
          end
        end

        aggregate
      end

      private

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

        if dev_or_test?
          raise ArgumentError, "Invalid call_site: #{@call_site.inspect}. Must match /\\A[a-z0-9_]+\\z/"
        else
        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(JSON.generate(
          severity: "WARN",
@@ -61,7 +61,6 @@ module Labkit
        ))
        @call_site = sanitized
      end
      end

      def rule_matches?(rule, identifier)
        rule.match.all? do |key, value|
@@ -76,7 +75,7 @@ module Labkit
          char_value = resolve_characteristic(char, @identifier)
          redis_key = build_redis_key(@call_site, index, char, char_value)

          count = increment_counter(redis_key, rule.period)
          count = incr_with_ttl(redis_key, rule.period)
          rule_exceeded = count > rule.limit

          exceeded = true if rule_exceeded
@@ -84,14 +83,13 @@ module Labkit
          log_rule(rule, index, count, redis_key, rule_exceeded)
        end

        (exceeded && rule.action == :block) ? :block : :allow
        exceeded && rule.action == :block ? :block : :allow
      end

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

          @logger.warn(JSON.generate(
            severity: "WARN",
            message: "rate_limit_unknown_characteristic",
@@ -99,7 +97,6 @@ module Labkit
          ))
          return UNKNOWN_SENTINEL
        end
        end

        value = identifier[char]

@@ -116,15 +113,15 @@ module Labkit

      def encode_char_value(value)
        if value.length > CHAR_VALUE_MAX_LENGTH
          Digest::SHA256.hexdigest(value)
          OpenSSL::Digest::SHA256.hexdigest(value)
        else
          value
        end
      end

      def increment_counter(redis_key, period)
      def incr_with_ttl(redis_key, period)
        count = @redis.incr(redis_key)
        # Set expiry only on first write (count == 1) to avoid resetting TTL on each call
        # Set expiry only on first write to avoid resetting TTL on each call
        @redis.expire(redis_key, period) if count == 1
        count
      end
@@ -159,7 +156,7 @@ module Labkit
      end

      def dev_or_test?
        env = ENV["LABKIT_ENV"]
        env = ENV.fetch("LABKIT_ENV", nil)
        env == "test" || env == "development"
      end

+7 −7
Original line number Diff line number Diff line
@@ -5,6 +5,13 @@ module Labkit
    # Identifier is a value object wrapping a hash of key-value pairs that
    # describe the caller (e.g. user, ip, endpoint).
    class Identifier
      # Normalize an endpoint value: strip query string.
      def self.normalize_endpoint(value)
        return value unless value.is_a?(String)

        value.split("?", 2).first
      end

      attr_reader :attributes

      def initialize(attributes = {})
@@ -24,13 +31,6 @@ module Labkit
      def ==(other)
        other.is_a?(Identifier) && other.attributes == @attributes
      end

      # Normalize an endpoint value: strip query string.
      def self.normalize_endpoint(value)
        return value unless value.is_a?(String)

        value.split("?", 2).first
      end
    end
  end
end
+3 −3
Original line number Diff line number Diff line
@@ -50,7 +50,7 @@ RSpec.describe Labkit::RateLimit::Evaluator do
      id = Labkit::RateLimit::Identifier.new(user: long_value)
      rule = make_rule(characteristics: [:user])

      expected_hash = Digest::SHA256.hexdigest(long_value)
      expected_hash = OpenSSL::Digest::SHA256.hexdigest(long_value)
      expect(redis).to receive(:incr).with("labkit:rl:rack_request:0:user:#{expected_hash}").and_return(1)
      expect(redis).to receive(:expire)

@@ -61,8 +61,8 @@ RSpec.describe Labkit::RateLimit::Evaluator do
      val_a = "a" + "x" * 200
      val_b = "b" + "x" * 200

      hash_a = Digest::SHA256.hexdigest(val_a)
      hash_b = Digest::SHA256.hexdigest(val_b)
      hash_a = OpenSSL::Digest::SHA256.hexdigest(val_a)
      hash_b = OpenSSL::Digest::SHA256.hexdigest(val_b)
      expect(hash_a).not_to eq(hash_b)
    end

+3 −3
Original line number Diff line number Diff line
@@ -263,7 +263,7 @@ RSpec.describe Labkit::RateLimit do

      check(identifier: id, rules: rules)

      expected_hash = Digest::SHA256.hexdigest(long_val)
      expected_hash = OpenSSL::Digest::SHA256.hexdigest(long_val)
      expect(redis.get("labkit:rl:rack_request:0:user:#{expected_hash}")).to eq(1)
    end

@@ -271,8 +271,8 @@ RSpec.describe Labkit::RateLimit do
      val_a = "prefix_" + "a" * 200
      val_b = "prefix_" + "b" * 200

      hash_a = Digest::SHA256.hexdigest(val_a)
      hash_b = Digest::SHA256.hexdigest(val_b)
      hash_a = OpenSSL::Digest::SHA256.hexdigest(val_a)
      hash_b = OpenSSL::Digest::SHA256.hexdigest(val_b)

      expect(hash_a).not_to eq(hash_b)
    end