/* SPDX-License-Identifier: GPL-3.0-only */ #include "ssh_auth_policy.h" #include #include #include #include #define CHECK(condition) do { \ if (!(condition)) { \ fprintf(stderr, "%s:%d: %s\n", __FILE__, __LINE__, #condition); \ exit(EXIT_FAILURE); \ } \ } while (0) _Static_assert(sizeof(ssh_auth_policy_t) <= 72, "policy must remain small"); _Static_assert(SSH_AUTH_POLICY_HANDSHAKE_CAPACITY == 6, "handshake burst"); _Static_assert(SSH_AUTH_POLICY_VERIFICATION_CAPACITY == 6, "verification burst"); _Static_assert(SSH_AUTH_POLICY_PROBE_CAPACITY == 12, "probe burst"); _Static_assert(SSH_AUTH_POLICY_HANDSHAKE_REFILL_US == 10000000, "handshake rate"); _Static_assert(SSH_AUTH_POLICY_VERIFICATION_REFILL_US == 10000000, "verification rate"); _Static_assert(SSH_AUTH_POLICY_PROBE_REFILL_US == 5000000, "probe rate"); static void drain(ssh_auth_policy_t *p, ssh_auth_policy_kind_t kind, int64_t now, unsigned count) { for (unsigned i = 0; i < count; ++i) { CHECK(ssh_auth_policy_admit(p, kind, now)); } CHECK(!ssh_auth_policy_admit(p, kind, now)); } static void unchanged(ssh_auth_policy_t *p, ssh_auth_policy_kind_t kind, int64_t now) { unsigned char before[sizeof(*p)]; memcpy(before, p, sizeof(*p)); CHECK(!ssh_auth_policy_admit(p, kind, now)); CHECK(memcmp(before, p, sizeof(*p)) == 0); } static void test_class(ssh_auth_policy_kind_t kind, unsigned capacity, int64_t interval) { ssh_auth_policy_t p = {0}; drain(&p, kind, 0, capacity); unchanged(&p, kind, -1); /* Initial zero is a real timestamp. */ CHECK(!ssh_auth_policy_admit(&p, kind, 1)); unchanged(&p, kind, 0); /* Regression after an empty-bucket denial. */ CHECK(!ssh_auth_policy_admit(&p, kind, interval - 1)); unchanged(&p, kind, interval - 2); CHECK(ssh_auth_policy_admit(&p, kind, interval)); CHECK(!ssh_auth_policy_admit(&p, kind, interval + 1)); CHECK(ssh_auth_policy_admit(&p, kind, 2 * interval)); /* Many denials cannot extend cooldown; exactly one token each interval. */ for (int64_t n = 3; n <= 1000; ++n) { CHECK(!ssh_auth_policy_admit(&p, kind, n * interval - 1)); CHECK(ssh_auth_policy_admit(&p, kind, n * interval)); CHECK(!ssh_auth_policy_admit(&p, kind, n * interval)); CHECK(!ssh_auth_policy_admit(&p, kind, n * interval + 1)); } p = (ssh_auth_policy_t){0}; drain(&p, kind, 0, capacity); /* Refill multiple tokens without losing fractional elapsed credit. */ CHECK(ssh_auth_policy_admit(&p, kind, 2 * interval + interval / 2)); CHECK(ssh_auth_policy_admit(&p, kind, 3 * interval - 1)); CHECK(!ssh_auth_policy_admit(&p, kind, 3 * interval - 1)); CHECK(ssh_auth_policy_admit(&p, kind, 3 * interval)); p = (ssh_auth_policy_t){0}; CHECK(ssh_auth_policy_admit(&p, kind, 0)); const int64_t idle = 100 * interval + interval / 2; drain(&p, kind, idle, capacity); /* Full idle discards fractional surplus. */ CHECK(!ssh_auth_policy_admit(&p, kind, idle + interval - 1)); CHECK(ssh_auth_policy_admit(&p, kind, idle + interval)); CHECK(!ssh_auth_policy_admit(&p, kind, idle + interval + 1)); p = (ssh_auth_policy_t){0}; drain(&p, kind, 0, capacity); const int64_t late = INT64_MAX - interval; drain(&p, kind, late, capacity); /* Huge forward jump saturates, not wraps. */ CHECK(!ssh_auth_policy_admit(&p, kind, INT64_MAX - 1)); CHECK(ssh_auth_policy_admit(&p, kind, INT64_MAX)); CHECK(!ssh_auth_policy_admit(&p, kind, INT64_MAX)); unchanged(&p, kind, INT64_MAX - 1); p = (ssh_auth_policy_t){0}; drain(&p, kind, INT64_MAX, capacity); /* Lazy init at maximum timestamp. */ p = (ssh_auth_policy_t){0}; drain(&p, kind, 0, capacity); drain(&p, kind, INT64_MAX, capacity); /* Direct maximum-sized subtraction. */ } static bool caller_a(ssh_auth_policy_t *p) { return ssh_auth_policy_admit(p, SSH_AUTH_POLICY_HANDSHAKE, 0); } static bool caller_b(ssh_auth_policy_t *p) { return ssh_auth_policy_admit(p, SSH_AUTH_POLICY_HANDSHAKE, 0); } static void test_validation_and_sharing(void) { ssh_auth_policy_t p = {0}; CHECK(!ssh_auth_policy_admit(NULL, SSH_AUTH_POLICY_HANDSHAKE, 0)); unchanged(&p, SSH_AUTH_POLICY_HANDSHAKE, INT64_MIN); unchanged(&p, (ssh_auth_policy_kind_t)-1, 0); unchanged(&p, SSH_AUTH_POLICY_KIND_COUNT, 0); unchanged(&p, (ssh_auth_policy_kind_t)INT_MAX, INT64_MAX); for (unsigned i = 0; i < 3; ++i) { CHECK(caller_a(&p)); CHECK(caller_b(&p)); } CHECK(!caller_a(&p)); CHECK(!caller_b(&p)); drain(&p, SSH_AUTH_POLICY_VERIFICATION, 0, 6); drain(&p, SSH_AUTH_POLICY_PROBE, 0, 12); CHECK(ssh_auth_policy_admit(&p, SSH_AUTH_POLICY_PROBE, 5000000)); CHECK(!ssh_auth_policy_admit(&p, SSH_AUTH_POLICY_HANDSHAKE, 5000000)); CHECK(!ssh_auth_policy_admit(&p, SSH_AUTH_POLICY_VERIFICATION, 5000000)); CHECK(ssh_auth_policy_admit(&p, SSH_AUTH_POLICY_HANDSHAKE, 10000000)); CHECK(ssh_auth_policy_admit(&p, SSH_AUTH_POLICY_VERIFICATION, 10000000)); CHECK(ssh_auth_policy_admit(&p, SSH_AUTH_POLICY_PROBE, 10000000)); unchanged(&p, (ssh_auth_policy_kind_t)-1, INT64_MAX); unchanged(&p, SSH_AUTH_POLICY_PROBE, -1); unchanged(&p, SSH_AUTH_POLICY_VERIFICATION, 9999999); p = (ssh_auth_policy_t){0}; CHECK(ssh_auth_policy_admit(&p, SSH_AUTH_POLICY_HANDSHAKE, 100)); unchanged(&p, SSH_AUTH_POLICY_PROBE, 99); /* Even an uninitialized class. */ drain(&p, SSH_AUTH_POLICY_PROBE, 100, 12); } int main(void) { test_class(SSH_AUTH_POLICY_HANDSHAKE, SSH_AUTH_POLICY_HANDSHAKE_CAPACITY, SSH_AUTH_POLICY_HANDSHAKE_REFILL_US); test_class(SSH_AUTH_POLICY_VERIFICATION, SSH_AUTH_POLICY_VERIFICATION_CAPACITY, SSH_AUTH_POLICY_VERIFICATION_REFILL_US); test_class(SSH_AUTH_POLICY_PROBE, SSH_AUTH_POLICY_PROBE_CAPACITY, SSH_AUTH_POLICY_PROBE_REFILL_US); test_validation_and_sharing(); printf("ssh_auth_policy: all tests passed (policy size: %zu bytes)\n", sizeof(ssh_auth_policy_t)); return EXIT_SUCCESS; }