/* SPDX-License-Identifier: GPL-3.0-only */ #include "ssh_auth_policy.h" #include bool ssh_auth_policy_admit(ssh_auth_policy_t *policy, ssh_auth_policy_kind_t kind, int64_t now_us) { if (policy == NULL || (unsigned)kind >= SSH_AUTH_POLICY_KIND_COUNT || now_us < 0) { return false; } for (unsigned i = 0; i < SSH_AUTH_POLICY_KIND_COUNT; ++i) { if (policy->buckets[i].initialized && now_us < policy->buckets[i].last_seen_us) { return false; } } const unsigned capacity = kind == SSH_AUTH_POLICY_PROBE ? SSH_AUTH_POLICY_PROBE_CAPACITY : kind == SSH_AUTH_POLICY_HANDSHAKE ? SSH_AUTH_POLICY_HANDSHAKE_CAPACITY : SSH_AUTH_POLICY_VERIFICATION_CAPACITY; const int64_t interval = kind == SSH_AUTH_POLICY_PROBE ? SSH_AUTH_POLICY_PROBE_REFILL_US : kind == SSH_AUTH_POLICY_HANDSHAKE ? SSH_AUTH_POLICY_HANDSHAKE_REFILL_US : SSH_AUTH_POLICY_VERIFICATION_REFILL_US; ssh_auth_policy_bucket_t *bucket = &policy->buckets[kind]; if (!bucket->initialized) { bucket->tokens = (uint8_t)capacity; bucket->refill_us = now_us; bucket->initialized = true; } else { /* Both timestamps are nonnegative and ordered. Divide before adding * to avoid overflow even for a jump from zero to INT64_MAX. */ const int64_t elapsed = now_us - bucket->refill_us; const int64_t earned = elapsed / interval; if (earned >= (int64_t)(capacity - bucket->tokens)) { bucket->tokens = (uint8_t)capacity; bucket->refill_us = now_us; } else { bucket->tokens += (uint8_t)earned; bucket->refill_us = now_us - elapsed % interval; } } bucket->last_seen_us = now_us; if (bucket->tokens == 0) { return false; } --bucket->tokens; return true; }