From a6431c0d6f468a03b661cc432a9b4205754ef0b4 Mon Sep 17 00:00:00 2001 From: Jonathan Payne Date: Tue, 29 Sep 2026 23:14:08 -0400 Subject: [PATCH 1/8] OpenConceptLab/ocl_online#275 | Capacity limit on heavy calls (semantic $match, $rerank), in shadow mode Counts semantic and reranked $match calls and $rerank calls in flight, in Redis lanes shared by the API tasks: cluster-wide, per task, semantic kNN, per plan tier and per user, with one cluster slot kept for single-row $match calls. Each call takes a lease in every lane that applies in one Lua script, renews it from a background thread and releases it when it finishes; a lease that isn't renewed expires, so a worker that dies frees its slots. - Shadow mode (the default) never refuses. Every gated call writes one JSON log line (decision, tier, lane, in-flight counts, endpoint, rows, held and queue time) for CloudWatch metric filters. - Responses carry X-OCL-Capacity-Decision, -Limit, -In-Flight, -Tier, -Tier-Limit, -Tier-In-Flight and -Suggested-Concurrency, exposed to browsers through CORS. - Enforce mode is built and tested, and ships off: a 429 with Retry-After before any work or quota charge. With enforce_for=aware, only clients that send capacity_aware in X-OCL-Event-Metadata are refused. - The mode and every number are runtime config: staff GET/PATCH/PUT /capacity/config/ or `manage.py capacity`, applied within 10 s, each change kept with who, when, old and new, and logged. - Fail open: if Redis can't be reached, calls go ahead uncounted and it's logged. The limiter uses its own short-timeout Redis client and skips Redis for 30 s after an error. - Tests run the real Lua scripts with fakeredis (CI has no Redis), and pass against Redis 7.0 too. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TyAC9YAn5hP3yGSREkSN8m --- README.md | 8 + core/capacity/__init__.py | 0 core/capacity/config.py | 194 ++++ core/capacity/constants.py | 56 ++ core/capacity/limiter.py | 505 +++++++++++ core/capacity/logs.py | 10 + core/capacity/management/__init__.py | 0 core/capacity/management/commands/__init__.py | 0 core/capacity/management/commands/capacity.py | 82 ++ core/capacity/migrations/0001_initial.py | 32 + core/capacity/migrations/__init__.py | 0 core/capacity/models.py | 22 + core/capacity/tests/__init__.py | 0 core/capacity/tests/tests.py | 841 ++++++++++++++++++ core/capacity/urls.py | 9 + core/capacity/views.py | 93 ++ core/common/utils.py | 12 + core/concepts/views.py | 57 +- core/settings.py | 16 + core/urls.py | 1 + core/users/constants.py | 1 + requirements.txt | 2 + 22 files changed, 1918 insertions(+), 23 deletions(-) create mode 100644 core/capacity/__init__.py create mode 100644 core/capacity/config.py create mode 100644 core/capacity/constants.py create mode 100644 core/capacity/limiter.py create mode 100644 core/capacity/logs.py create mode 100644 core/capacity/management/__init__.py create mode 100644 core/capacity/management/commands/__init__.py create mode 100644 core/capacity/management/commands/capacity.py create mode 100644 core/capacity/migrations/0001_initial.py create mode 100644 core/capacity/migrations/__init__.py create mode 100644 core/capacity/models.py create mode 100644 core/capacity/tests/__init__.py create mode 100644 core/capacity/tests/tests.py create mode 100644 core/capacity/urls.py create mode 100644 core/capacity/views.py diff --git a/README.md b/README.md index a06c6a5a8..180c6e6c4 100644 --- a/README.md +++ b/README.md @@ -36,6 +36,14 @@ API supports the OpenID implicit flow. If `OIDC_SERVER_URL` and `OIDC_REALM` are not provided then the Django Auth is enabled by default. +#### Capacity limit on heavy calls +Semantic `$match` calls (kNN search and/or the in-request rerank) and `$rerank` calls are heavy: each holds an API worker for seconds. The capacity limit counts how many run at once in Redis lanes (cluster-wide, per API process host, semantic kNN, per plan tier and per user) and sends `X-OCL-Capacity-*` headers on those responses. It has three modes: +- `off`: nothing is counted. +- `shadow` (the default): calls are counted and logged, and none is refused. +- `enforce`: a call that finds a lane full gets a 429 with `Retry-After`, before any work or quota charge. With `enforce_for=aware` (the default) only clients that send `"capacity_aware": "true"` in `X-OCL-Event-Metadata` are refused; the rest stay in shadow mode. + +The mode and all the numbers are runtime settings, changed by staff with `GET`/`PATCH`/`PUT /capacity/config/` or `python manage.py capacity show|set|reset|history|status`. Each change is kept in `/capacity/config/history/`, and every API process picks it up within `CAPACITY_CONFIG_CACHE_SECONDS` (10). `CAPACITY_LIMIT_MODE` sets the mode until staff first change it. If Redis can't be reached, calls go ahead uncounted. + ### Run Checks (use the `docker exec` command in a service started with `docker compose up -d`) 1. Pylint (pep8): diff --git a/core/capacity/__init__.py b/core/capacity/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/core/capacity/config.py b/core/capacity/config.py new file mode 100644 index 000000000..58181530b --- /dev/null +++ b/core/capacity/config.py @@ -0,0 +1,194 @@ +""" +The capacity limit's settings: the mode and every number, changed at runtime by staff through +`/capacity/config/` or `manage.py capacity`, with no deploy or restart. Each change adds a CapacityConfig row +(who, when, the old config and the new one) and writes one log line. API processes re-read the newest row at +most every CAPACITY_CONFIG_CACHE_SECONDS, so a change applies across the cluster within that time. +""" +import copy +import logging +import threading +import time + +from django.conf import settings +from django.core.exceptions import ValidationError +from django.db import connection, transaction + +from core.capacity.constants import ( + MODES, MODE_SHADOW, ENFORCE_FOR, ENFORCE_FOR_AWARE, TIER_STAFF, TIER_CORE, TIER_EARLY_ACCESS, + TIER_PREVIEW, CONFIG_LOG_EVENT, SOURCE_API) +from core.capacity.logs import emit + +logger = logging.getLogger('oclapi') + +# The starting numbers (OpenConceptLab/ocl_online#275). Counts are heavy calls in flight at once. +DEFAULTS = { + 'mode': MODE_SHADOW, # the boot default comes from settings.CAPACITY_LIMIT_MODE + 'enforce_for': ENFORCE_FOR_AWARE, + 'api_heavy': { + 'cluster': 4, # across all API tasks: half their workers, so the rest stay free for everything else + 'per_task': 2, # on each API task, since the load balancer doesn't see how busy a task is + }, + 'es2_knn': 3, # semantic $match calls, which run kNN searches in Elasticsearch + 'reserve_single_row': 1, # cluster slots only single-row $match calls may use, so interactive calls don't wait + 'tiers': {TIER_STAFF: 4, TIER_CORE: 4, TIER_EARLY_ACCESS: 3, TIER_PREVIEW: 2}, # ceilings, not reservations + 'per_user': {TIER_STAFF: 4, TIER_CORE: 3, TIER_EARLY_ACCESS: 2, TIER_PREVIEW: 1}, + 'lease_seconds': 60, # a call's lease expires this long after its last renewal, so a dead worker's frees + 'renew_seconds': 20, + 'max_hold_seconds': 900, # stop renewing after this, in case a call never finishes + 'retry_after': { + 'base': 5, # seconds per call ahead in the fullest lane + 'max': 30, + 'paused': 120, # when a lane's limit is 0 + }, +} +CHOICES = {'mode': MODES, 'enforce_for': ENFORCE_FOR} +MAX_NUMBER = 100000 +ADVISORY_LOCK_ID = 275275 # serializes config changes, so a PATCH never merges onto a stale version + +_cache = {'config': None, 'expires_at': 0.0} +_cache_lock = threading.Lock() + + +def get_defaults(): + defaults = copy.deepcopy(DEFAULTS) + mode = getattr(settings, 'CAPACITY_LIMIT_MODE', None) + if mode in MODES: + defaults['mode'] = mode + return defaults + + +def merge(base, changes, known_only=False): + """`base` with `changes` merged in, dict by dict. With known_only, keys `base` doesn't have are dropped.""" + result = copy.deepcopy(base) + for key, value in (changes or {}).items(): + if known_only and key not in base: + continue + if isinstance(value, dict) and isinstance(result.get(key), dict): + result[key] = merge(result[key], value, known_only) + else: + result[key] = copy.deepcopy(value) + return result + + +def diff(old, new, prefix=''): + """{dotted.path: [old, new]} for every value that differs.""" + changes = {} + for key in sorted(set(old or {}) | set(new or {})): + before, after = (old or {}).get(key), (new or {}).get(key) + path = f'{prefix}{key}' + if isinstance(before, dict) and isinstance(after, dict): + changes.update(diff(before, after, f'{path}.')) + elif before != after: + changes[path] = [before, after] + return changes + + +def _validate_shape(config, template, prefix, errors): + if not isinstance(config, dict): + errors.append(f'{prefix.rstrip(".") or "config"} must be an object.') + return + for key in config: + if key not in template: + errors.append(f'Unknown setting "{prefix}{key}".') + for key, default in template.items(): + path = f'{prefix}{key}' + if key not in config: + errors.append(f'"{path}" is required.') + elif isinstance(default, dict): + _validate_shape(config[key], default, f'{path}.', errors) + elif key in CHOICES: + if config[key] not in CHOICES[key]: + errors.append(f'"{path}" must be one of {", ".join(CHOICES[key])}.') + elif isinstance(config[key], bool) or not isinstance(config[key], int) or not ( + 0 <= config[key] <= MAX_NUMBER): + errors.append(f'"{path}" must be a whole number from 0 to {MAX_NUMBER}.') + + +def validate(config): + """Raise ValidationError unless `config` is a complete, consistent config.""" + errors = [] + _validate_shape(config, DEFAULTS, '', errors) + if not errors: + if config['reserve_single_row'] > config['api_heavy']['cluster']: + errors.append('"reserve_single_row" can\'t be more than "api_heavy.cluster".') + if config['lease_seconds'] < 5: + errors.append('"lease_seconds" must be at least 5.') + if not 1 <= config['renew_seconds'] < config['lease_seconds']: + errors.append('"renew_seconds" must be at least 1 and less than "lease_seconds".') + if config['max_hold_seconds'] < config['lease_seconds']: + errors.append('"max_hold_seconds" can\'t be less than "lease_seconds".') + retry_after = config['retry_after'] + if retry_after['base'] < 1 or retry_after['paused'] < 1 or retry_after['max'] < retry_after['base']: + errors.append('"retry_after" needs "base" and "paused" of at least 1, and "max" of at least "base".') + if errors: + raise ValidationError(errors) + return config + + +def resolve(stored): + """The config a stored version puts in force: the defaults, overlaid with what it sets.""" + config = merge(get_defaults(), stored, known_only=True) + try: + return validate(config) + except ValidationError as ex: + logger.error('Capacity config is invalid (%s); using the defaults', '; '.join(ex.messages)) + return get_defaults() + + +def get_current(): + """(the newest CapacityConfig row or None, the config in force), read from the database.""" + from core.capacity.models import CapacityConfig + latest = CapacityConfig.get_latest() + return latest, resolve(latest.config if latest else None) + + +def get_config(): + """The config in force, cached per process for CAPACITY_CONFIG_CACHE_SECONDS. Never raises.""" + now = time.monotonic() + config = _cache['config'] + if config is not None and now < _cache['expires_at']: + return config + with _cache_lock: + if _cache['config'] is not None and time.monotonic() < _cache['expires_at']: + return _cache['config'] # another thread just refreshed it + try: + _, config = get_current() + except Exception as ex: + # Keep the last good config (or the defaults) rather than fail the request. + logger.warning('Capacity config could not be read (%s); using the last known config', ex) + config = _cache['config'] or get_defaults() + _cache['config'] = config + _cache['expires_at'] = now + settings.CAPACITY_CONFIG_CACHE_SECONDS + return config + + +def clear_cache(): + with _cache_lock: + _cache['config'] = None + _cache['expires_at'] = 0.0 + + +def save_config(changes, user=None, source=SOURCE_API, note='', replace=False): + """ + Put a new version in force: `changes` merged onto the config in force, or onto the defaults with replace. + Raises ValidationError if the result isn't valid. Returns (row, config, {path: [old, new]}); when nothing + changes, no row is added and row is the current one (or None). + """ + from core.capacity.models import CapacityConfig + with transaction.atomic(): + if connection.vendor == 'postgresql': + with connection.cursor() as cursor: + cursor.execute('SELECT pg_advisory_xact_lock(%s)', [ADVISORY_LOCK_ID]) + latest, previous = get_current() + config = validate(merge(get_defaults() if replace else previous, changes)) + changed = diff(previous, config) + if not changed: + return latest, previous, {} + row = CapacityConfig.objects.create( + config=config, previous_config=previous, created_by=user, source=source, note=note or '') + clear_cache() + emit({ + 'event': CONFIG_LOG_EVENT, 'version': row.id, 'changed_by': getattr(user, 'username', None), + 'source': source, 'note': note or None, 'changes': changed, 'mode': config['mode'], + }) + return row, config, changed diff --git a/core/capacity/constants.py b/core/capacity/constants.py new file mode 100644 index 000000000..248ae015f --- /dev/null +++ b/core/capacity/constants.py @@ -0,0 +1,56 @@ +MODE_OFF = 'off' # no counting, no headers +MODE_SHADOW = 'shadow' # count, log and send headers, but never refuse +MODE_ENFORCE = 'enforce' # refuse with 429 + Retry-After when a lane is full +MODES = (MODE_OFF, MODE_SHADOW, MODE_ENFORCE) + +# Which clients enforce mode refuses. The rest stay in shadow: counted and logged, never refused. +ENFORCE_FOR_AWARE = 'aware' # only clients whose X-OCL-Event-Metadata carries "capacity_aware": "true" +ENFORCE_FOR_ALL = 'all' +ENFORCE_FOR = (ENFORCE_FOR_AWARE, ENFORCE_FOR_ALL) +CAPACITY_AWARE_METADATA_KEY = 'capacity_aware' + +# A user's tier is their highest plan group, highest first. +TIER_STAFF = 'staff' +TIER_CORE = 'core' +TIER_EARLY_ACCESS = 'early_access' +TIER_PREVIEW = 'preview' +TIERS = (TIER_STAFF, TIER_CORE, TIER_EARLY_ACCESS, TIER_PREVIEW) + +# Lanes: each counts the heavy calls in flight in one scope. +LANE_API_HEAVY = 'api_heavy' # every heavy call, cluster-wide +LANE_API_HEAVY_TASK = 'api_heavy_task' # every heavy call on this API task +LANE_ES2_KNN = 'es2_knn' # semantic $match calls, which run kNN searches in Elasticsearch +LANE_TIER = 'tier' # the calls of every user in this tier +LANE_USER = 'user' # this user's calls +LANES = (LANE_API_HEAVY, LANE_API_HEAVY_TASK, LANE_ES2_KNN, LANE_TIER, LANE_USER) + +DECISION_ADMITTED = 'admitted' +DECISION_SHADOW_REFUSED = 'shadow-refused' # a lane was full; enforce mode would have refused it +DECISION_REFUSED = 'refused' +DECISION_UNAVAILABLE = 'unavailable' # Redis couldn't be reached, so the call went ahead uncounted + +ENDPOINT_MATCH = '$match' +ENDPOINT_RERANK = '$rerank' + +HEADER_DECISION = 'X-OCL-Capacity-Decision' +HEADER_LIMIT = 'X-OCL-Capacity-Limit' +HEADER_IN_FLIGHT = 'X-OCL-Capacity-In-Flight' +HEADER_TIER = 'X-OCL-Capacity-Tier' +HEADER_TIER_LIMIT = 'X-OCL-Capacity-Tier-Limit' +HEADER_TIER_IN_FLIGHT = 'X-OCL-Capacity-Tier-In-Flight' +HEADER_SUGGESTED_CONCURRENCY = 'X-OCL-Capacity-Suggested-Concurrency' +HEADERS = ( + HEADER_DECISION, HEADER_LIMIT, HEADER_IN_FLIGHT, HEADER_TIER, HEADER_TIER_LIMIT, HEADER_TIER_IN_FLIGHT, + HEADER_SUGGESTED_CONCURRENCY, +) + +CAPACITY_EXCEEDED_ERROR_CODE = 'capacity_exceeded' + +# The "event" of the JSON log lines, which CloudWatch metric filters match on. +LOG_EVENT = 'ocl_capacity' +CONFIG_LOG_EVENT = 'ocl_capacity_config' + +REDIS_KEY_PREFIX = 'ocl:capacity' + +SOURCE_API = 'api' +SOURCE_COMMAND = 'command' diff --git a/core/capacity/limiter.py b/core/capacity/limiter.py new file mode 100644 index 000000000..d0332f02b --- /dev/null +++ b/core/capacity/limiter.py @@ -0,0 +1,505 @@ +""" +The capacity limit on heavy calls (OpenConceptLab/ocl_online#275): semantic `$match` (kNN searches and/or the +in-request rerank) and `$rerank`. It caps how many run at once across all users, to protect the API workers and +Elasticsearch. It isn't a quota: it charges nothing, and a call it refuses (in enforce mode only) is asked to retry shortly. + +Each lane is a Redis sorted set of leases: member = a random token per call, score = the lease's expiry in ms by +Redis's own clock, so the API tasks' clocks don't matter. A call takes a lease in every lane that applies to it, in +one atomic script, renews them from a background thread while it runs, and releases them when it finishes. A lease +that isn't renewed expires, so a worker that dies frees its slots within `lease_seconds`. + +If Redis can't be reached, calls go ahead uncounted and it's logged (fail open): the limiter must never cause an +outage. It uses its own Redis client with short timeouts and no retries, and skips Redis for +CAPACITY_REDIS_RETRY_SECONDS after an error, so an outage costs a call at most one short timeout. +""" +import socket +import threading +import time +import uuid +from contextlib import contextmanager + +import redis +from cid.locals import get_cid +from django.conf import settings +from redis.backoff import NoBackoff +from redis.retry import Retry +from redis.sentinel import Sentinel +from rest_framework import status +from rest_framework.response import Response + +from core.capacity.config import get_config +from core.capacity.constants import ( + TIERS, MODE_OFF, MODE_ENFORCE, ENFORCE_FOR_ALL, CAPACITY_AWARE_METADATA_KEY, TIER_STAFF, TIER_CORE, + TIER_EARLY_ACCESS, TIER_PREVIEW, LANE_API_HEAVY, LANE_API_HEAVY_TASK, LANE_ES2_KNN, LANE_TIER, LANE_USER, + DECISION_ADMITTED, DECISION_SHADOW_REFUSED, DECISION_REFUSED, DECISION_UNAVAILABLE, ENDPOINT_MATCH, + HEADER_DECISION, HEADER_LIMIT, HEADER_IN_FLIGHT, HEADER_TIER, HEADER_TIER_LIMIT, HEADER_TIER_IN_FLIGHT, + HEADER_SUGGESTED_CONCURRENCY, CAPACITY_EXCEEDED_ERROR_CODE, LOG_EVENT, REDIS_KEY_PREFIX) +from core.capacity.logs import emit +from core.common.utils import get_event_metadata +from core.users.constants import CORE_USER_GROUP, EARLY_ACCESS_GROUP + +_NOW_MS = """ +if redis.replicate_commands then redis.replicate_commands() end +local now = redis.call('TIME') +local now_ms = tonumber(now[1]) * 1000 + math.floor(tonumber(now[2]) / 1000) +""" + +# KEYS: one per lane. ARGV: token, lease ms, force ('1': take the lease even when a lane is full), then one limit +# per key. Returns {1 if the lease was taken else 0, then each lane's count before this call}. +ACQUIRE_SCRIPT = _NOW_MS + """ +local lease_ms = tonumber(ARGV[2]) +local result = {0} +local full = false +for i, key in ipairs(KEYS) do + redis.call('ZREMRANGEBYSCORE', key, '-inf', now_ms) + local count = redis.call('ZCARD', key) + result[i + 1] = count + if count >= tonumber(ARGV[3 + i]) then full = true end +end +if ARGV[3] == '1' or not full then + result[1] = 1 + for _, key in ipairs(KEYS) do + redis.call('ZADD', key, now_ms + lease_ms, ARGV[1]) + redis.call('PEXPIRE', key, lease_ms * 2) + end +end +return result +""" + +# KEYS: the call's lanes. ARGV: token, lease ms. Extends the leases the call still holds; returns how many. +RENEW_SCRIPT = _NOW_MS + """ +local lease_ms = tonumber(ARGV[2]) +local renewed = 0 +for _, key in ipairs(KEYS) do + if redis.call('ZSCORE', key, ARGV[1]) then + redis.call('ZADD', key, now_ms + lease_ms, ARGV[1]) + redis.call('PEXPIRE', key, lease_ms * 2) + renewed = renewed + 1 + end +end +return renewed +""" + +# KEYS: lanes. Returns each lane's count of unexpired leases. +COUNT_SCRIPT = _NOW_MS + """ +local counts = {} +for i, key in ipairs(KEYS) do + counts[i] = redis.call('ZCOUNT', key, '(' .. now_ms, '+inf') +end +return counts +""" + + +def lane_key(lane, *parts): + return ':'.join([REDIS_KEY_PREFIX, lane, *[str(part) for part in parts]]) + + +def get_task_id(): + """Names this API task's lane. The API tasks run on separate hosts, so the hostname tells them apart.""" + return settings.CAPACITY_TASK_ID or socket.gethostname() + + +def build_redis_client(): + timeout = settings.CAPACITY_REDIS_TIMEOUT_SECONDS + options = { + 'socket_timeout': timeout, 'socket_connect_timeout': timeout, 'retry': Retry(NoBackoff(), 0), + 'retry_on_timeout': False, 'health_check_interval': 0, + } + if settings.REDIS_SENTINELS: + sentinel_options = {'socket_timeout': timeout, 'socket_connect_timeout': timeout} + if settings.REDIS_PASSWORD: + sentinel_options['password'] = settings.REDIS_PASSWORD + sentinel = Sentinel(settings.REDIS_SENTINELS_LIST, sentinel_kwargs=sentinel_options) + return sentinel.master_for( + settings.REDIS_SENTINELS_MASTER, db=settings.REDIS_DB, password=settings.REDIS_PASSWORD, **options) + return redis.Redis( + host=settings.REDIS_HOST, port=int(settings.REDIS_PORT), db=settings.REDIS_DB, + password=settings.REDIS_PASSWORD, **options) + + +class RedisLanes: + """The lanes in Redis, with a per-process switch that skips Redis for a while after it fails.""" + client = None + scripts = {} + unavailable_until = 0.0 + lock = threading.Lock() + + @classmethod + def use_client(cls, client): + with cls.lock: + cls.client = client + cls.scripts = {} + cls.unavailable_until = 0.0 + + @classmethod + def get_client(cls): + if cls.client is None: + with cls.lock: + if cls.client is None: + cls.client = build_redis_client() + cls.scripts = {} + return cls.client + + @classmethod + def run_script(cls, source, keys, args=()): + client = cls.get_client() + script = cls.scripts.get(source) + if script is None or script.registered_client is not client: + script = cls.scripts[source] = client.register_script(source) + return script(keys=keys, args=args) + + @classmethod + def is_available(cls): + return time.monotonic() >= cls.unavailable_until + + @classmethod + def mark_unavailable(cls): + cls.unavailable_until = time.monotonic() + settings.CAPACITY_REDIS_RETRY_SECONDS + + @classmethod + def acquire(cls, keys, limits, token, lease_ms, force): # pylint: disable=too-many-arguments + """(whether the lease was taken, each lane's count before this call)""" + result = cls.run_script(ACQUIRE_SCRIPT, keys, [token, lease_ms, '1' if force else '0', *limits]) + return bool(result[0]), [int(count) for count in result[1:]] + + @classmethod + def renew(cls, keys, token, lease_ms): + return int(cls.run_script(RENEW_SCRIPT, keys, [token, lease_ms])) + + @classmethod + def release(cls, keys, token): + pipeline = cls.get_client().pipeline(transaction=False) + for key in keys: + pipeline.zrem(key, token) + pipeline.execute() + + @classmethod + def count(cls, keys): + return [int(count) for count in cls.run_script(COUNT_SCRIPT, keys)] if keys else [] + + @classmethod + def find_keys(cls, lane): + return sorted(key.decode() if isinstance(key, bytes) else key + for key in cls.get_client().scan_iter(match=lane_key(lane, '*'), count=100)) + + +def get_tier(user): + """The user's highest plan tier, as capabilities resolve it: staff > core > early_access > preview.""" + if user.is_staff or user.is_superuser: + return TIER_STAFF + groups = set(user.groups.values_list('name', flat=True)) + if CORE_USER_GROUP in groups: + return TIER_CORE + if EARLY_ACCESS_GROUP in groups: + return TIER_EARLY_ACCESS + return TIER_PREVIEW + + +def get_queue_ms(request, now=None): + """ + How long the request waited before the view started, when an AWS load balancer stamped it. The load balancer + adds X-Amzn-Trace-Id with the time it received the request, in whole seconds (in `Self`, or in `Root` when it + started the trace), so this reads up to a second high. Mostly it's time spent queued for a free API worker. + """ + fields = dict(part.split('=', 1) for part in (request.META.get('HTTP_X_AMZN_TRACE_ID') or '').split(';') + if '=' in part) + try: + received = int((fields.get('Self') or fields.get('Root')).split('-')[1], 16) + except (AttributeError, IndexError, ValueError): + return None + waited = (now or time.time()) - received + return max(int(waited * 1000), 0) if -2 < waited < 3600 else None + + +class LeaseRenewer(threading.Thread): + """Renews a call's leases every `renew_seconds` until it's stopped, or for `max_hold_seconds` at most.""" + def __init__(self, gate): + super().__init__(name='ocl-capacity-lease', daemon=True) + self.gate = gate + self.interval = gate.config['renew_seconds'] + self.deadline = time.monotonic() + gate.config['max_hold_seconds'] + self.finished = threading.Event() + + def run(self): + while not self.finished.wait(self.interval) and time.monotonic() < self.deadline: + self.gate.renew() + + def stop(self): + self.finished.set() + self.join(timeout=2 * settings.CAPACITY_REDIS_TIMEOUT_SECONDS + 1) + + +class CapacityGate: + """ + Admission for one heavy call. `acquire()` takes a lease in each lane that applies, `release()` gives it back and + writes the call's log line. In shadow mode (and in enforce mode for clients it doesn't enforce for), a call that + finds a lane full still goes ahead, and the line says it would have been refused. Never raises. + """ + def __init__(self, request, endpoint, rows=0, semantic=False, reranker=False): # pylint: disable=too-many-arguments + self.request = request + self.endpoint = endpoint + self.rows = rows + self.semantic = semantic + self.reranker = reranker + self.single_row = endpoint == ENDPOINT_MATCH and rows == 1 + self.token = uuid.uuid4().hex + self.config = None + self.tier = None + self.lanes = [] # [(lane, Redis key, limit)] + self.counts = {} # lane: calls in flight, counting this one if it holds a lease + self.full = [] # lanes that were already at their limit + self.decision = None # None: the limit is off, or the call isn't gated + self.enforced = False + self.holding = False + self.retry_after = None + self.queue_ms = None + self.error = None + self.lease_lost = False + self.renewer = None + self.started_at = None + self.released_at = None + + @property + def refused(self): + return self.decision == DECISION_REFUSED + + @property + def keys(self): + return [key for _, key, _ in self.lanes] + + @property + def limits(self): + return {lane: limit for lane, _, limit in self.lanes} + + @property + def lease_ms(self): + return self.config['lease_seconds'] * 1000 + + @property + def metadata(self): + return get_event_metadata(self.request) + + def acquire(self): + self.started_at = time.monotonic() + try: + self.config = get_config() + if self.config['mode'] == MODE_OFF: + return self + self.queue_ms = get_queue_ms(self.request) + self.tier = get_tier(self.request.user) + self.lanes = self.get_lanes() + self.enforced = self.is_enforced() + if not RedisLanes.is_available(): + self.decision, self.error = DECISION_UNAVAILABLE, 'skipped: Redis failed recently' + return self + taken, before = RedisLanes.acquire( + self.keys, [limit for _, _, limit in self.lanes], self.token, self.lease_ms, force=not self.enforced) + self.holding = taken + self.full = [lane for (lane, _, limit), count in zip(self.lanes, before) if count >= limit] + self.counts = {lane: count + int(taken) for (lane, _, _), count in zip(self.lanes, before)} + if not self.full: + self.decision = DECISION_ADMITTED + else: + self.decision = DECISION_SHADOW_REFUSED if taken else DECISION_REFUSED + self.retry_after = self.get_retry_after(dict(zip([lane for lane, _, _ in self.lanes], before))) + if taken: + self.renewer = LeaseRenewer(self) + self.renewer.start() + except Exception as ex: + self.record_error(ex) + self.decision = self.decision or DECISION_UNAVAILABLE + return self + + def get_lanes(self): + config = self.config + cluster = config['api_heavy']['cluster'] + if not self.single_row: + cluster = max(cluster - config['reserve_single_row'], 0) + lanes = [ + (LANE_API_HEAVY, lane_key(LANE_API_HEAVY), cluster), + (LANE_API_HEAVY_TASK, lane_key(LANE_API_HEAVY_TASK, get_task_id()), config['api_heavy']['per_task']), + ] + if self.semantic: + lanes.append((LANE_ES2_KNN, lane_key(LANE_ES2_KNN), config['es2_knn'])) + lanes.append((LANE_TIER, lane_key(LANE_TIER, self.tier), config['tiers'][self.tier])) + lanes.append((LANE_USER, lane_key(LANE_USER, self.request.user.id), config['per_user'][self.tier])) + return lanes + + def is_enforced(self): + if self.config['mode'] != MODE_ENFORCE: + return False + if self.config['enforce_for'] == ENFORCE_FOR_ALL: + return True + return str(self.metadata.get(CAPACITY_AWARE_METADATA_KEY)).lower() in ('true', '1') + + def get_retry_after(self, before): + retry_after, limits = self.config['retry_after'], self.limits + if any(limits[lane] == 0 for lane in self.full): + return retry_after['paused'] + ahead = max(before[lane] - limits[lane] + 1 for lane in self.full) + return min(retry_after['max'], retry_after['base'] * ahead) + + def get_suggested_concurrency(self): + """How many calls this user should keep in flight right now: 0 while a lane is paused, else at least 1.""" + limits = self.limits + if not limits: + return None + if 0 in limits.values(): + return 0 + if not self.counts: + return limits[LANE_USER] + headroom = min(limit - self.counts[lane] for lane, limit in limits.items() if lane != LANE_API_HEAVY_TASK) + return max(1, min(limits[LANE_USER], self.counts[LANE_USER] + headroom)) + + def renew(self): + try: + if RedisLanes.renew(self.keys, self.token, self.lease_ms) < len(self.lanes): + self.lease_lost = True + except Exception as ex: + self.record_error(ex) + + def release(self): + if self.released_at is not None: + return + self.released_at = time.monotonic() + try: + if self.renewer: + self.renewer.stop() + if self.holding: + RedisLanes.release(self.keys, self.token) + except Exception as ex: + self.record_error(ex) + finally: + if self.decision: + try: + self.log() + except Exception: # a log line must never fail the call + pass + + def record_error(self, ex): + if isinstance(ex, redis.RedisError): + RedisLanes.mark_unavailable() + self.error = f'{ex.__class__.__name__}: {ex}'[:300] + + def get_refusal_response(self): + scope = self.full[0] + return Response( + { + 'detail': f'Matching is busy. Please retry in {self.retry_after} seconds.', + 'error_code': CAPACITY_EXCEEDED_ERROR_CODE, + 'scope': scope, + 'paused': self.limits[scope] == 0, + 'retry_after': self.retry_after, + }, + status=status.HTTP_429_TOO_MANY_REQUESTS + ) + + def apply_headers(self, response): + if not self.decision or response is None: + return + try: + limits, counts = self.limits, self.counts + headers = { + HEADER_DECISION: self.decision, + HEADER_LIMIT: limits.get(LANE_API_HEAVY), + HEADER_IN_FLIGHT: counts.get(LANE_API_HEAVY), + HEADER_TIER: self.tier, + HEADER_TIER_LIMIT: limits.get(LANE_TIER), + HEADER_TIER_IN_FLIGHT: counts.get(LANE_TIER), + HEADER_SUGGESTED_CONCURRENCY: self.get_suggested_concurrency(), + } + if self.refused: + headers['Retry-After'] = self.retry_after + for name, value in headers.items(): + if value is not None: + response[name] = str(value) + except Exception as ex: + self.record_error(ex) + + def log(self): + limits, counts, metadata = self.limits, self.counts, self.metadata + + def lane_fields(lane, prefix): + return {f'{prefix}in_flight': counts.get(lane), f'{prefix}limit': limits.get(lane)} + + emit({ + 'event': LOG_EVENT, + 'decision': self.decision, + 'mode': self.config['mode'], + 'enforced': self.enforced, + 'endpoint': self.endpoint, + 'semantic': self.semantic, + 'reranker': self.reranker, + 'rows': self.rows, + 'single_row': self.single_row, + 'tier': self.tier, + 'user_id': self.request.user.id, + 'scope': self.full[0] if self.full else None, + 'full': self.full or None, + **lane_fields(LANE_API_HEAVY, ''), + **lane_fields(LANE_API_HEAVY_TASK, 'task_'), + **lane_fields(LANE_ES2_KNN, 'es2_knn_'), + **lane_fields(LANE_TIER, 'tier_'), + **lane_fields(LANE_USER, 'user_'), + 'suggested': self.get_suggested_concurrency(), + 'retry_after': self.retry_after, + 'held_ms': int((self.released_at - self.started_at) * 1000) if self.holding else None, + 'queue_ms': self.queue_ms, + 'lease_lost': self.lease_lost or None, + 'error': self.error, + 'task': get_task_id(), + 'request_source': str(self.request.META.get('HTTP_X_OCL_REQUEST_SOURCE') or '')[:50] or None, + 'algorithm_id': str(metadata.get('algorithm_id') or '')[:100] or None, + 'automatch_run_id': str(metadata.get('automatch_run_id') or '')[:20] or None, + 'cid': get_cid(), + }) + + +class CapacityLimitMixin: + """For views with heavy calls: gate a call with `capacity_gate`, and every response gets the capacity headers.""" + capacity_gate_instance = None + + @contextmanager + def capacity_gate(self, request, **call): + gate = self.capacity_gate_instance = CapacityGate(request, **call).acquire() + try: + yield gate + finally: + gate.release() + + def finalize_response(self, request, response, *args, **kwargs): + response = super().finalize_response(request, response, *args, **kwargs) + if self.capacity_gate_instance is not None: + self.capacity_gate_instance.apply_headers(response) + return response + + +def get_status(config=None): + """What's in flight in each lane now, against its limit. For staff; reads Redis directly.""" + config = config or get_config() + status_ = {'mode': config['mode'], 'enforce_for': config['enforce_for']} + try: + task_keys = RedisLanes.find_keys(LANE_API_HEAVY_TASK) + user_keys = RedisLanes.find_keys(LANE_USER) + tier_keys = [lane_key(LANE_TIER, tier) for tier in TIERS] + keys = [lane_key(LANE_API_HEAVY), lane_key(LANE_ES2_KNN), *tier_keys, *task_keys, *user_keys] + counts = dict(zip(keys, RedisLanes.count(keys))) + except Exception as ex: + status_['redis'] = f'{ex.__class__.__name__}: {ex}'[:300] + return status_ + + def suffix(key): + return key.rsplit(':', 1)[-1] + + status_.update({ + 'redis': 'ok', + LANE_API_HEAVY: {'in_flight': counts[lane_key(LANE_API_HEAVY)], 'limit': config['api_heavy']['cluster'], + 'bulk_limit': config['api_heavy']['cluster'] - config['reserve_single_row']}, + LANE_ES2_KNN: {'in_flight': counts[lane_key(LANE_ES2_KNN)], 'limit': config['es2_knn']}, + 'tasks': {suffix(key): {'in_flight': counts[key], 'limit': config['api_heavy']['per_task']} + for key in task_keys if counts[key]}, + 'tiers': {tier: {'in_flight': counts[lane_key(LANE_TIER, tier)], 'limit': config['tiers'][tier]} + for tier in TIERS}, + 'users': {suffix(key): {'in_flight': counts[key]} for key in user_keys if counts[key]}, + }) + return status_ diff --git a/core/capacity/logs.py b/core/capacity/logs.py new file mode 100644 index 000000000..841d32f65 --- /dev/null +++ b/core/capacity/logs.py @@ -0,0 +1,10 @@ +import json + + +def emit(record): + """ + Write one JSON object as one line of the API log, where CloudWatch metric filters match it. The API's other + timing lines are prints too: gunicorn captures stdout, and the `oclapi` logger has no handler in production. + """ + print(json.dumps({key: value for key, value in record.items() if value is not None}, + separators=(',', ':'), default=str), flush=True) diff --git a/core/capacity/management/__init__.py b/core/capacity/management/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/core/capacity/management/commands/__init__.py b/core/capacity/management/commands/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/core/capacity/management/commands/capacity.py b/core/capacity/management/commands/capacity.py new file mode 100644 index 000000000..8c633758f --- /dev/null +++ b/core/capacity/management/commands/capacity.py @@ -0,0 +1,82 @@ +import json + +from django.core.exceptions import ValidationError +from django.core.management import BaseCommand, CommandError + +from core.capacity.config import get_current, save_config, diff +from core.capacity.constants import SOURCE_COMMAND +from core.capacity.limiter import get_status + +# manage.py capacity show the config in force +# manage.py capacity set mode=enforce tiers.preview=1 [--note ...] [--user ] +# manage.py capacity reset [--note ...] back to the defaults +# manage.py capacity history [--limit N] changes, newest first +# manage.py capacity status heavy calls in flight now, by lane + + +def parse_assignment(assignment): + """'tiers.preview=1' -> {'tiers': {'preview': 1}}""" + path, sep, raw = assignment.partition('=') + if not sep or not path: + raise CommandError(f'Expected =, got "{assignment}".') + value = int(raw) if raw.lstrip('-').isdigit() else raw + for key in reversed(path.split('.')): + value = {key: value} + return value + + +class Command(BaseCommand): + help = 'Show or change the capacity limit on heavy calls (semantic $match, $rerank).' + + def add_arguments(self, parser): + parser.add_argument('action', choices=['show', 'set', 'reset', 'history', 'status']) + parser.add_argument('assignments', nargs='*', help='For set: =, e.g. tiers.preview=1') + parser.add_argument('--note', default='', help='Why, recorded with the change') + parser.add_argument('--user', default=None, help='Username to record the change under') + parser.add_argument('--limit', type=int, default=20, help='For history: how many changes') + + def handle(self, *args, **options): + action = options['action'] + if action == 'show': + row, config = get_current() + self.print({'version': row.id if row else None, 'config': config}) + elif action == 'status': + self.print(get_status(get_current()[1])) + elif action == 'history': + from core.capacity.models import CapacityConfig + for row in CapacityConfig.objects.select_related('created_by').order_by('-id')[:options['limit']]: + self.print({ + 'version': row.id, 'created_at': row.created_at.isoformat(), 'source': row.source, + 'created_by': row.created_by.username if row.created_by else None, 'note': row.note or None, + 'changes': diff(row.previous_config, row.config), + }) + else: + self.save(action, options) + + def save(self, action, options): + if action == 'set' and not options['assignments']: + raise CommandError('Give at least one =.') + changes = {} + for assignment in options['assignments'] if action == 'set' else []: + changes = self.merge(changes, parse_assignment(assignment)) + user = None + if options['user']: + from core.users.models import UserProfile + user = UserProfile.objects.filter(username=options['user']).first() + if not user: + raise CommandError(f'No user "{options["user"]}".') + try: + row, _, changed = save_config( + changes, user=user, source=SOURCE_COMMAND, note=options['note'], replace=action == 'reset') + except ValidationError as ex: + raise CommandError(' '.join(ex.messages)) from ex + self.print({'version': row.id if row else None, 'changes': changed}) + + @classmethod + def merge(cls, base, changes): + for key, value in changes.items(): + base[key] = cls.merge(base.get(key, {}), value) if isinstance(value, dict) else value + return base + + def print(self, data): + self.stdout.write(json.dumps(data, indent=2, default=str)) diff --git a/core/capacity/migrations/0001_initial.py b/core/capacity/migrations/0001_initial.py new file mode 100644 index 000000000..20b8a6165 --- /dev/null +++ b/core/capacity/migrations/0001_initial.py @@ -0,0 +1,32 @@ +# Generated by Django 5.2.17 on 2026-09-30 02:58 + +import django.db.models.deletion +from django.conf import settings +from django.db import migrations, models + + +class Migration(migrations.Migration): + + initial = True + + dependencies = [ + migrations.swappable_dependency(settings.AUTH_USER_MODEL), + ] + + operations = [ + migrations.CreateModel( + name='CapacityConfig', + fields=[ + ('id', models.AutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')), + ('config', models.JSONField()), + ('previous_config', models.JSONField(blank=True, null=True)), + ('created_at', models.DateTimeField(auto_now_add=True)), + ('source', models.CharField(max_length=16)), + ('note', models.TextField(blank=True, default='')), + ('created_by', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='+', to=settings.AUTH_USER_MODEL)), + ], + options={ + 'db_table': 'capacity_configs', + }, + ), + ] diff --git a/core/capacity/migrations/__init__.py b/core/capacity/migrations/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/core/capacity/models.py b/core/capacity/models.py new file mode 100644 index 000000000..9ec566485 --- /dev/null +++ b/core/capacity/models.py @@ -0,0 +1,22 @@ +from django.db import models + + +class CapacityConfig(models.Model): + """ + One version of the capacity-limit settings (core.capacity.config). Rows are only ever added: the newest is the + config in force, and the older ones are its history. With no row at all, the defaults apply. + """ + class Meta: + db_table = 'capacity_configs' + + config = models.JSONField() # the full config this version put in force + previous_config = models.JSONField(null=True, blank=True) # the config in force before it + created_by = models.ForeignKey( + 'users.UserProfile', null=True, blank=True, on_delete=models.SET_NULL, related_name='+') + created_at = models.DateTimeField(auto_now_add=True) + source = models.CharField(max_length=16) # 'api' or 'command' + note = models.TextField(blank=True, default='') + + @classmethod + def get_latest(cls): + return cls.objects.order_by('-id').first() diff --git a/core/capacity/tests/__init__.py b/core/capacity/tests/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/core/capacity/tests/tests.py b/core/capacity/tests/tests.py new file mode 100644 index 000000000..550c82bc4 --- /dev/null +++ b/core/capacity/tests/tests.py @@ -0,0 +1,841 @@ +import json +import os +import time +from io import StringIO +from unittest.mock import patch, Mock + +import fakeredis +import redis +from django.conf import settings +from django.contrib.auth.models import Group +from django.core.exceptions import ValidationError +from django.core.management import call_command, CommandError +from django.test import RequestFactory, override_settings + +from core.capacity.config import ( + get_defaults, validate, merge, diff, resolve, save_config, get_config, clear_cache, _cache) +from core.capacity.constants import ( + HEADERS, HEADER_DECISION, HEADER_LIMIT, HEADER_IN_FLIGHT, HEADER_TIER, HEADER_TIER_LIMIT, HEADER_TIER_IN_FLIGHT, + HEADER_SUGGESTED_CONCURRENCY, ENDPOINT_MATCH, ENDPOINT_RERANK, DECISION_ADMITTED, DECISION_SHADOW_REFUSED, + DECISION_REFUSED, DECISION_UNAVAILABLE, LANE_API_HEAVY, LANE_API_HEAVY_TASK, LANE_ES2_KNN, LANE_TIER, LANE_USER, + LOG_EVENT, CONFIG_LOG_EVENT, SOURCE_COMMAND) +from core.capacity.limiter import ( + RedisLanes, CapacityGate, LeaseRenewer, get_tier, get_queue_ms, lane_key, get_task_id, build_redis_client, + get_status) +from core.capacity.models import CapacityConfig +from core.common.tests import OCLTestCase, OCLAPITestCase, PREVIEW_GROUP_NAME +from core.common.utils import get_event_metadata +from core.users.constants import CORE_USER_GROUP, EARLY_ACCESS_GROUP +from core.users.tests.factories import UserProfileFactory + + +def make_user(group=None, **kwargs): + user = UserProfileFactory(**kwargs) + if group: + user.groups.add(Group.objects.get_or_create(name=group)[0]) + return user + + +class CapacityTestMixin: + """Lanes in a fake Redis (fakeredis runs the real Lua scripts), a fresh config cache, and captured log lines.""" + def setUp(self): + super().setUp() + url = os.environ.get('CAPACITY_TEST_REDIS_URL') # set it to run these against a real (throwaway) Redis + self.redis = redis.Redis.from_url(url) if url else fakeredis.FakeRedis(server=fakeredis.FakeServer()) + if url: + self.redis.flushdb() + RedisLanes.use_client(self.redis) + clear_cache() + self.limiter_emit = patch('core.capacity.limiter.emit').start() + self.config_emit = patch('core.capacity.config.emit').start() + self.gates = [] + + def tearDown(self): + for gate in self.gates: + gate.release() + patch.stopall() + RedisLanes.use_client(None) + clear_cache() + super().tearDown() + + @staticmethod + def configure(**changes): + return save_config(changes) + + def acquire( # pylint: disable=too-many-arguments + self, user, endpoint=ENDPOINT_MATCH, rows=10, semantic=True, reranker=False, metadata=None, **meta): + extra = {'HTTP_X_OCL_EVENT_METADATA': json.dumps(metadata)} if metadata else {} + request = RequestFactory().post('/concepts/$match/', **extra, **meta) + request.user = user + gate = CapacityGate(request, endpoint=endpoint, rows=rows, semantic=semantic, reranker=reranker).acquire() + self.gates.append(gate) + return gate + + def lines(self): + return [call.args[0] for call in self.limiter_emit.call_args_list] + + +class CapacityConfigTest(OCLTestCase): + def setUp(self): + super().setUp() + clear_cache() + self.emit = patch('core.capacity.config.emit').start() + + def tearDown(self): + patch.stopall() + clear_cache() + super().tearDown() + + def test_defaults_are_the_starting_numbers(self): + defaults = validate(get_defaults()) + + self.assertEqual(defaults['mode'], 'shadow') + self.assertEqual(defaults['api_heavy'], {'cluster': 4, 'per_task': 2}) + self.assertEqual(defaults['es2_knn'], 3) + self.assertEqual(defaults['tiers'], {'staff': 4, 'core': 4, 'early_access': 3, 'preview': 2}) + self.assertEqual(defaults['per_user'], {'staff': 4, 'core': 3, 'early_access': 2, 'preview': 1}) + self.assertEqual(defaults['reserve_single_row'], 1) + + def test_default_mode_comes_from_settings(self): + with override_settings(CAPACITY_LIMIT_MODE='off'): + self.assertEqual(get_defaults()['mode'], 'off') + with override_settings(CAPACITY_LIMIT_MODE='loud'): + self.assertEqual(get_defaults()['mode'], 'shadow') + + def test_validate_rejects_bad_configs(self): + cases = [ + ({'bogus': 1}, 'Unknown setting "bogus".'), + ({'tiers': {'gold': 1}}, 'Unknown setting "tiers.gold".'), + ({'tiers': 5}, 'tiers must be an object.'), + ({'mode': 'loud'}, '"mode" must be one of off, shadow, enforce.'), + ({'enforce_for': 'some'}, '"enforce_for" must be one of aware, all.'), + ({'es2_knn': -1}, '"es2_knn" must be a whole number'), + ({'es2_knn': True}, '"es2_knn" must be a whole number'), + ({'es2_knn': '3'}, '"es2_knn" must be a whole number'), + ({'per_user': {'preview': 100001}}, '"per_user.preview" must be a whole number'), + ({'reserve_single_row': 5}, '"reserve_single_row" can\'t be more than "api_heavy.cluster".'), + ({'lease_seconds': 4, 'renew_seconds': 1}, '"lease_seconds" must be at least 5.'), + ({'renew_seconds': 60}, '"renew_seconds" must be at least 1 and less than "lease_seconds".'), + ({'renew_seconds': 0}, '"renew_seconds" must be at least 1 and less than "lease_seconds".'), + ({'max_hold_seconds': 30}, '"max_hold_seconds" can\'t be less than "lease_seconds".'), + ({'retry_after': {'base': 0}}, '"retry_after" needs'), + ({'retry_after': {'base': 10, 'max': 5}}, '"retry_after" needs'), + ] + for changes, message in cases: + with self.subTest(changes=changes): + with self.assertRaises(ValidationError) as context: + validate(merge(get_defaults(), changes)) + self.assertIn(message, ' '.join(context.exception.messages)) + + with self.assertRaises(ValidationError) as context: + validate({}) + self.assertIn('"mode" is required.', context.exception.messages) + with self.assertRaises(ValidationError): + validate([]) + + def test_merge_and_diff(self): + base = {'a': 1, 'b': {'c': 2, 'd': 3}} + + self.assertEqual(merge(base, {'b': {'c': 5}, 'e': 6}), {'a': 1, 'b': {'c': 5, 'd': 3}, 'e': 6}) + self.assertEqual(merge(base, {'b': {'c': 5}, 'e': 6}, known_only=True), {'a': 1, 'b': {'c': 5, 'd': 3}}) + self.assertEqual(base, {'a': 1, 'b': {'c': 2, 'd': 3}}) + self.assertEqual(diff(base, {'a': 1, 'b': {'c': 5, 'd': 3}, 'e': 6}), {'b.c': [2, 5], 'e': [None, 6]}) + self.assertEqual(diff(None, {'a': 1}), {'a': [None, 1]}) + self.assertEqual(diff(base, base), {}) + + def test_resolve(self): + self.assertEqual(resolve(None), get_defaults()) + self.assertEqual(resolve({'es2_knn': 5, 'retired_setting': 1})['es2_knn'], 5) + with self.assertLogs('oclapi', level='ERROR'): + self.assertEqual(resolve({'es2_knn': -1}), get_defaults()) + + def test_save_config_adds_a_version_with_its_history(self): + user = UserProfileFactory() + + row, config, changes = save_config({'tiers': {'preview': 1}, 'mode': 'off'}, user=user, note='quiet') + + self.assertEqual(changes, {'mode': ['shadow', 'off'], 'tiers.preview': [2, 1]}) + self.assertEqual(config['tiers']['preview'], 1) + self.assertEqual(row.config, config) + self.assertEqual(row.previous_config, get_defaults()) + self.assertEqual(row.created_by, user) + self.assertEqual((row.source, row.note), ('api', 'quiet')) + self.assertEqual(CapacityConfig.get_latest(), row) + self.assertEqual(get_config(), config) + self.emit.assert_called_once_with({ + 'event': CONFIG_LOG_EVENT, 'version': row.id, 'changed_by': user.username, 'source': 'api', + 'note': 'quiet', 'changes': changes, 'mode': 'off', + }) + + row2, config2, changes2 = save_config({'es2_knn': 2}, source=SOURCE_COMMAND) + self.assertEqual(changes2, {'es2_knn': [3, 2]}) + self.assertEqual(row2.previous_config, config) + self.assertEqual(config2['tiers']['preview'], 1) + self.assertIsNone(row2.created_by) + + same_row, same_config, no_changes = save_config({'es2_knn': 2}) + self.assertEqual((same_row, same_config, no_changes), (row2, config2, {})) + self.assertEqual(CapacityConfig.objects.count(), 2) + + def test_save_config_replace_starts_from_the_defaults(self): + save_config({'tiers': {'preview': 1}, 'es2_knn': 1}) + + _, config, changes = save_config({'mode': 'off'}, replace=True) + + self.assertEqual(config, {**get_defaults(), 'mode': 'off'}) + self.assertEqual(changes, {'es2_knn': [1, 3], 'mode': ['shadow', 'off'], 'tiers.preview': [1, 2]}) + + def test_save_config_rejects_an_invalid_change(self): + with self.assertRaises(ValidationError): + save_config({'es2_knn': -1}) + self.assertFalse(CapacityConfig.objects.exists()) + self.emit.assert_not_called() + + @override_settings(CAPACITY_CONFIG_CACHE_SECONDS=60) + def test_get_config_is_cached(self): + self.assertEqual(get_config()['es2_knn'], 3) + CapacityConfig.objects.create(config={**get_defaults(), 'es2_knn': 1}, source='api') + + self.assertEqual(get_config()['es2_knn'], 3) + clear_cache() + self.assertEqual(get_config()['es2_knn'], 1) + + def test_get_config_keeps_the_last_known_config_when_the_database_fails(self): + with patch('core.capacity.config.get_current', side_effect=Exception('db down')): + with self.assertLogs('oclapi', level='WARNING'): + self.assertEqual(get_config(), get_defaults()) + + save_config({'es2_knn': 1}) + get_config() + _cache['expires_at'] = 0.0 + with patch('core.capacity.config.get_current', side_effect=Exception('db down')): + with self.assertLogs('oclapi', level='WARNING'): + self.assertEqual(get_config()['es2_knn'], 1) + + +class CapacityHelpersTest(OCLTestCase): + def test_get_tier(self): + self.assertEqual(get_tier(make_user(is_staff=True)), 'staff') + self.assertEqual(get_tier(make_user(is_superuser=True)), 'staff') + self.assertEqual(get_tier(make_user(CORE_USER_GROUP)), 'core') + self.assertEqual(get_tier(make_user(EARLY_ACCESS_GROUP)), 'early_access') + self.assertEqual(get_tier(make_user(PREVIEW_GROUP_NAME)), 'preview') + self.assertEqual(get_tier(make_user()), 'preview') + user = make_user(PREVIEW_GROUP_NAME) + user.groups.add(Group.objects.get(name=CORE_USER_GROUP)) + self.assertEqual(get_tier(user), 'core') + + def test_get_queue_ms(self): + def request(trace_id=None): + return RequestFactory().post('/', **({'HTTP_X_AMZN_TRACE_ID': trace_id} if trace_id else {})) + + received = 1790000000 + root = f'Root=1-{received:x}-abcdef012345678912345678' + self.assertEqual(get_queue_ms(request(root), now=received + 2.5), 2500) + self.assertEqual( + get_queue_ms(request(f'Root=1-{received - 100:x}-abc;Self=1-{received:x}-def'), now=received + 1), 1000) + self.assertEqual(get_queue_ms(request(root), now=received - 1), 0) + self.assertIsNone(get_queue_ms(request(root), now=received + 7200)) + self.assertIsNone(get_queue_ms(request(), now=received)) + self.assertIsNone(get_queue_ms(request('Root=1-zz-abc'), now=received)) + self.assertIsNone(get_queue_ms(request('Root=nope'), now=received)) + self.assertIsNone(get_queue_ms(request('garbage'), now=received)) + + def test_get_event_metadata(self): + def request(value=None): + return RequestFactory().post('/', **({'HTTP_X_OCL_EVENT_METADATA': value} if value else {})) + + self.assertEqual(get_event_metadata(request()), {}) + self.assertEqual(get_event_metadata(request('{nope')), {}) + self.assertEqual(get_event_metadata(request('[1]')), {}) + self.assertEqual( + get_event_metadata(request('{"algorithm_id": "ocl-semantic"}')), {'algorithm_id': 'ocl-semantic'}) + + def test_lane_key_and_task_id(self): + self.assertEqual(lane_key(LANE_API_HEAVY), 'ocl:capacity:api_heavy') + self.assertEqual(lane_key(LANE_USER, 42), 'ocl:capacity:user:42') + with override_settings(CAPACITY_TASK_ID='task-a'): + self.assertEqual(get_task_id(), 'task-a') + with override_settings(CAPACITY_TASK_ID=''): + with patch('core.capacity.limiter.socket.gethostname', return_value='h'): + self.assertEqual(get_task_id(), 'h') + + @override_settings(REDIS_SENTINELS=None, REDIS_HOST='redis-host', REDIS_PORT='6390', REDIS_PASSWORD=None, + CAPACITY_REDIS_TIMEOUT_SECONDS=0.25) + def test_build_redis_client(self): + client = build_redis_client() + + kwargs = client.connection_pool.connection_kwargs + self.assertEqual((kwargs['host'], kwargs['port']), ('redis-host', 6390)) + self.assertEqual((kwargs['socket_timeout'], kwargs['socket_connect_timeout']), (0.25, 0.25)) + self.assertEqual(kwargs['retry']._retries, 0) # pylint: disable=protected-access + + @override_settings( + REDIS_SENTINELS='s1:26379;s2:26379', REDIS_SENTINELS_LIST=[('s1', 26379), ('s2', 26379)], + REDIS_SENTINELS_MASTER='primary', REDIS_PASSWORD='secret', CAPACITY_REDIS_TIMEOUT_SECONDS=0.25) + def test_build_redis_client_with_sentinels(self): + with patch('core.capacity.limiter.Sentinel') as sentinel_mock: + client = build_redis_client() + + self.assertEqual(client, sentinel_mock.return_value.master_for.return_value) + sentinel_mock.assert_called_once_with( + [('s1', 26379), ('s2', 26379)], + sentinel_kwargs={'socket_timeout': 0.25, 'socket_connect_timeout': 0.25, 'password': 'secret'}) + args, kwargs = sentinel_mock.return_value.master_for.call_args + self.assertEqual(args, ('primary',)) + self.assertEqual((kwargs['password'], kwargs['socket_timeout']), ('secret', 0.25)) + + +class CapacityGateTest(CapacityTestMixin, OCLTestCase): + def test_admits_and_counts_a_call_in_every_lane(self): + user = make_user(EARLY_ACCESS_GROUP) + + gate = self.acquire(user) + + self.assertEqual(gate.decision, DECISION_ADMITTED) + self.assertTrue(gate.holding) + self.assertEqual(gate.tier, 'early_access') + self.assertEqual(gate.limits, { + LANE_API_HEAVY: 3, LANE_API_HEAVY_TASK: 2, LANE_ES2_KNN: 3, LANE_TIER: 3, LANE_USER: 2}) + self.assertEqual(gate.counts, { + LANE_API_HEAVY: 1, LANE_API_HEAVY_TASK: 1, LANE_ES2_KNN: 1, LANE_TIER: 1, LANE_USER: 1}) + for key in gate.keys: + self.assertEqual(self.redis.zcard(key), 1) + self.assertGreater(self.redis.pttl(key), 0) + + gate.release() + for key in gate.keys: + self.assertEqual(self.redis.zcard(key), 0) + gate.release() # releasing twice is harmless + self.assertEqual(len(self.lines()), 1) + + def test_shadow_mode_never_refuses_but_says_it_would_have(self): + user = make_user(PREVIEW_GROUP_NAME) + + first = self.acquire(user) + second = self.acquire(user) + + self.assertEqual(first.decision, DECISION_ADMITTED) + self.assertEqual(second.decision, DECISION_SHADOW_REFUSED) + self.assertTrue(second.holding) + self.assertEqual(second.full, [LANE_USER]) + self.assertEqual(second.counts[LANE_USER], 2) + self.assertEqual(second.retry_after, 5) # 5 s for each call ahead in the fullest lane + self.assertFalse(second.refused) + + def test_enforce_mode_refuses_without_taking_a_lease(self): + self.configure(mode='enforce', enforce_for='all') + user = make_user(PREVIEW_GROUP_NAME) + + first = self.acquire(user) + second = self.acquire(user) + + self.assertEqual(first.decision, DECISION_ADMITTED) + self.assertEqual(second.decision, DECISION_REFUSED) + self.assertTrue(second.refused) + self.assertFalse(second.holding) + self.assertIsNone(second.renewer) + self.assertEqual(second.counts[LANE_USER], 1) + self.assertEqual(self.redis.zcard(lane_key(LANE_USER, user.id)), 1) + response = second.get_refusal_response() + self.assertEqual(response.status_code, 429) + self.assertEqual(response.data, { + 'detail': 'Matching is busy. Please retry in 5 seconds.', 'error_code': 'capacity_exceeded', + 'scope': LANE_USER, 'paused': False, 'retry_after': 5, + }) + + first.release() + self.assertEqual(self.acquire(user).decision, DECISION_ADMITTED) + + def test_enforce_for_aware_refuses_only_clients_that_say_they_handle_it(self): + self.configure(mode='enforce') + user = make_user(PREVIEW_GROUP_NAME) + self.acquire(user) + + unaware = self.acquire(user, metadata={'algorithm_id': 'ocl-semantic'}) + aware = self.acquire(user, metadata={'capacity_aware': 'true'}) + + self.assertEqual(unaware.decision, DECISION_SHADOW_REFUSED) + self.assertFalse(unaware.enforced) + self.assertEqual(aware.decision, DECISION_REFUSED) + self.assertTrue(aware.enforced) + + def test_a_stale_lease_expires(self): + user = make_user(PREVIEW_GROUP_NAME) + # a lease left by a worker that died: its expiry has passed + self.redis.zadd(lane_key(LANE_USER, user.id), {'dead-worker': 1000}) + + gate = self.acquire(user) + + self.assertEqual(gate.decision, DECISION_ADMITTED) + self.assertEqual(self.redis.zrange(lane_key(LANE_USER, user.id), 0, -1), [gate.token.encode()]) + + def test_single_row_calls_can_use_the_reserved_slot(self): + self.configure(mode='enforce', enforce_for='all', api_heavy={'cluster': 2, 'per_task': 5}) + bulk = self.acquire(make_user(CORE_USER_GROUP)) + + second_bulk = self.acquire(make_user(CORE_USER_GROUP)) + single_row = self.acquire(make_user(CORE_USER_GROUP), rows=1) + rerank = self.acquire(make_user(CORE_USER_GROUP), endpoint=ENDPOINT_RERANK, rows=1, semantic=False) + + self.assertEqual(bulk.decision, DECISION_ADMITTED) + self.assertEqual(bulk.limits[LANE_API_HEAVY], 1) + self.assertEqual((second_bulk.decision, second_bulk.full), (DECISION_REFUSED, [LANE_API_HEAVY])) + self.assertTrue(single_row.single_row) + self.assertEqual(single_row.limits[LANE_API_HEAVY], 2) + self.assertEqual(single_row.decision, DECISION_ADMITTED) + self.assertFalse(rerank.single_row) # $rerank calls are part of Auto Match runs, never interactive + self.assertEqual(rerank.decision, DECISION_REFUSED) + + def test_each_api_task_has_its_own_lane(self): + self.configure(mode='enforce', enforce_for='all', api_heavy={'cluster': 4, 'per_task': 1}, reserve_single_row=0) + with override_settings(CAPACITY_TASK_ID='task-a'): + first = self.acquire(make_user(CORE_USER_GROUP)) + second = self.acquire(make_user(CORE_USER_GROUP)) + with override_settings(CAPACITY_TASK_ID='task-b'): + other_task = self.acquire(make_user(CORE_USER_GROUP)) + + self.assertEqual(first.decision, DECISION_ADMITTED) + self.assertEqual((second.decision, second.full), (DECISION_REFUSED, [LANE_API_HEAVY_TASK])) + self.assertEqual(other_task.decision, DECISION_ADMITTED) + self.assertEqual(other_task.counts[LANE_API_HEAVY], 2) + + def test_tier_ceilings_leave_other_tiers_alone(self): + self.configure(mode='enforce', enforce_for='all', api_heavy={'cluster': 10, 'per_task': 10}, es2_knn=10) + self.acquire(make_user(PREVIEW_GROUP_NAME)) + self.acquire(make_user(PREVIEW_GROUP_NAME)) + + third_preview = self.acquire(make_user(PREVIEW_GROUP_NAME)) + early_access = self.acquire(make_user(EARLY_ACCESS_GROUP)) + + self.assertEqual((third_preview.decision, third_preview.full), (DECISION_REFUSED, [LANE_TIER])) + self.assertEqual(third_preview.counts[LANE_TIER], 2) + self.assertEqual(early_access.decision, DECISION_ADMITTED) + self.assertEqual(early_access.counts[LANE_TIER], 1) + + def test_only_semantic_calls_take_an_es2_knn_lane(self): + self.configure(mode='enforce', enforce_for='all', es2_knn=1) + semantic = self.acquire(make_user(CORE_USER_GROUP)) + + semantic_again = self.acquire(make_user(CORE_USER_GROUP)) + reranked_lexical = self.acquire(make_user(CORE_USER_GROUP), semantic=False, reranker=True) + + self.assertIn(LANE_ES2_KNN, semantic.limits) + self.assertEqual((semantic_again.decision, semantic_again.full), (DECISION_REFUSED, [LANE_ES2_KNN])) + self.assertNotIn(LANE_ES2_KNN, reranked_lexical.limits) + self.assertEqual(reranked_lexical.decision, DECISION_ADMITTED) + + def test_a_paused_tier(self): + self.configure(mode='enforce', enforce_for='all', tiers={'preview': 0}) + + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + + self.assertEqual((gate.decision, gate.full), (DECISION_REFUSED, [LANE_TIER])) + self.assertEqual(gate.retry_after, 120) + self.assertEqual(gate.get_suggested_concurrency(), 0) + self.assertTrue(gate.get_refusal_response().data['paused']) + + def test_retry_after_grows_with_the_calls_ahead_up_to_the_max(self): + self.configure(tiers={'core': 10}, per_user={'core': 10}, api_heavy={'cluster': 10, 'per_task': 10}) + user = make_user(CORE_USER_GROUP) + gates = [self.acquire(user, semantic=False, reranker=True) for _ in range(12)] + + # bulk calls get 9 of the cluster's 10 + self.assertEqual([gate.retry_after for gate in gates[8:]], [None, 5, 10, 15]) + self.configure(retry_after={'base': 5, 'max': 12, 'paused': 120}) + self.assertEqual(self.acquire(user, semantic=False, reranker=True).retry_after, 12) + + def test_suggested_concurrency(self): + core = make_user(CORE_USER_GROUP) + + first = self.acquire(core) + self.assertEqual(first.get_suggested_concurrency(), 3) # alone: up to the per-user limit + second = self.acquire(make_user(CORE_USER_GROUP)) + self.assertEqual(second.get_suggested_concurrency(), 2) # its own + the 1 bulk slot left in the cluster + preview = self.acquire(make_user(PREVIEW_GROUP_NAME)) + self.assertEqual(preview.get_suggested_concurrency(), 1) # never below 1 unless paused + self.assertIsNone(CapacityGate(Mock(), ENDPOINT_MATCH).get_suggested_concurrency()) + + def test_renewal_extends_the_leases_it_still_holds(self): + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + key = lane_key(LANE_USER, gate.request.user.id) + self.redis.zadd(key, {gate.token: 1000}) + + gate.renew() + + self.assertGreater(self.redis.zscore(key, gate.token), time.time() * 1000) + self.assertFalse(gate.lease_lost) + self.redis.delete(key) + gate.renew() + self.assertTrue(gate.lease_lost) + self.assertFalse(self.redis.exists(key)) # a lost lease isn't re-added + + def test_the_renewer_thread_renews_until_stopped(self): + gate = Mock(config={'renew_seconds': 0.01, 'max_hold_seconds': 60}) + renewer = LeaseRenewer(gate) + renewer.start() + time.sleep(0.1) + renewer.stop() + + self.assertFalse(renewer.is_alive()) + self.assertGreater(gate.renew.call_count, 1) + + gate = Mock(config={'renew_seconds': 0.01, 'max_hold_seconds': 0}) + renewer = LeaseRenewer(gate) + renewer.start() + renewer.join(1) + self.assertFalse(renewer.is_alive()) # gave up at max_hold_seconds + gate.renew.assert_not_called() + + def test_a_call_starts_a_renewer_and_release_stops_it(self): + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + + self.assertTrue(gate.renewer.is_alive()) + gate.release() + self.assertFalse(gate.renewer.is_alive()) + + def test_redis_failure_fails_open_and_skips_redis_for_a_while(self): + broken = Mock() + broken.register_script.return_value = Mock(side_effect=redis.ConnectionError('refused')) + RedisLanes.use_client(broken) + + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + + self.assertEqual(gate.decision, DECISION_UNAVAILABLE) + self.assertFalse(gate.holding) + self.assertEqual(gate.error, 'ConnectionError: refused') + self.assertFalse(RedisLanes.is_available()) + + skipped = self.acquire(make_user(PREVIEW_GROUP_NAME)) + self.assertEqual(skipped.decision, DECISION_UNAVAILABLE) + self.assertEqual(broken.register_script.return_value.call_count, 1) + self.assertEqual(skipped.get_suggested_concurrency(), 1) + + RedisLanes.unavailable_until = 0.0 + self.acquire(make_user(PREVIEW_GROUP_NAME)) + self.assertEqual(broken.register_script.return_value.call_count, 2) + + def test_a_failed_release_or_renewal_is_logged_not_raised(self): + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + with patch.object(RedisLanes, 'renew', side_effect=redis.TimeoutError('slow')): + gate.renew() + self.assertEqual(gate.error, 'TimeoutError: slow') + + with patch.object(RedisLanes, 'release', side_effect=redis.ConnectionError('gone')): + gate.release() + self.assertEqual(self.lines()[-1]['error'], 'ConnectionError: gone') + self.assertFalse(RedisLanes.is_available()) + + def test_an_unexpected_error_fails_open(self): + with patch('core.capacity.limiter.get_tier', side_effect=KeyError('tier')): + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + + self.assertEqual(gate.decision, DECISION_UNAVAILABLE) + self.assertEqual(gate.error, "KeyError: 'tier'") + self.assertTrue(RedisLanes.is_available()) # only a Redis error skips Redis + gate.release() + self.assertEqual(self.lines()[-1]['decision'], DECISION_UNAVAILABLE) + + def test_off_mode_does_nothing(self): + self.configure(mode='off') + with patch.object(RedisLanes, 'acquire') as acquire_mock: + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + response = {} + + gate.apply_headers(response) + gate.release() + + acquire_mock.assert_not_called() + self.assertIsNone(gate.decision) + self.assertEqual(response, {}) + self.assertEqual(self.lines(), []) + + def test_headers(self): + self.configure(mode='enforce', enforce_for='all') + user = make_user(PREVIEW_GROUP_NAME) + admitted, refused = self.acquire(user), self.acquire(user) + admitted_headers, refused_headers = {}, {} + + admitted.apply_headers(admitted_headers) + refused.apply_headers(refused_headers) + + self.assertEqual(admitted_headers, { + HEADER_DECISION: 'admitted', HEADER_LIMIT: '3', HEADER_IN_FLIGHT: '1', HEADER_TIER: 'preview', + HEADER_TIER_LIMIT: '2', HEADER_TIER_IN_FLIGHT: '1', HEADER_SUGGESTED_CONCURRENCY: '1', + }) + self.assertEqual(refused_headers[HEADER_DECISION], 'refused') + self.assertEqual(refused_headers['Retry-After'], '5') + admitted.apply_headers(None) + broken = Mock() + broken.__setitem__ = Mock(side_effect=ValueError('bad header')) + admitted.apply_headers(broken) + self.assertEqual(admitted.error, 'ValueError: bad header') + + def test_the_log_line(self): + user = make_user(PREVIEW_GROUP_NAME) + self.acquire(user) + received = int(time.time()) - 3 + gate = self.acquire( + user, rows=25, reranker=True, metadata={'algorithm_id': 'ocl-semantic', 'automatch_run_id': '77'}, + HTTP_X_OCL_REQUEST_SOURCE='automatch', HTTP_X_AMZN_TRACE_ID=f'Root=1-{received:x}-abc') + + gate.release() + + line = self.lines()[-1] + self.assertGreaterEqual(line.pop('queue_ms'), 3000) + self.assertGreaterEqual(line.pop('held_ms'), 0) + line.pop('cid') + self.assertEqual(line, { + 'event': LOG_EVENT, 'decision': DECISION_SHADOW_REFUSED, 'mode': 'shadow', 'enforced': False, + 'endpoint': '$match', 'semantic': True, 'reranker': True, 'rows': 25, 'single_row': False, + 'tier': 'preview', 'user_id': user.id, 'scope': LANE_USER, 'full': [LANE_USER], + 'in_flight': 2, 'limit': 3, 'task_in_flight': 2, 'task_limit': 2, 'es2_knn_in_flight': 2, + 'es2_knn_limit': 3, 'tier_in_flight': 2, 'tier_limit': 2, 'user_in_flight': 2, 'user_limit': 1, + 'suggested': 1, 'retry_after': 5, 'lease_lost': None, 'error': None, 'task': get_task_id(), + 'request_source': 'automatch', 'algorithm_id': 'ocl-semantic', 'automatch_run_id': '77', + }) + + def test_a_failing_log_line_never_fails_the_call(self): + self.limiter_emit.side_effect = OSError('stdout closed') + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + + gate.release() + + self.assertEqual(self.limiter_emit.call_count, 1) + + def test_status(self): + with override_settings(CAPACITY_TASK_ID='task-a'): + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + self.acquire(make_user(CORE_USER_GROUP), semantic=False, reranker=True) + + state = get_status() + + self.assertEqual(state['redis'], 'ok') + self.assertEqual(state[LANE_API_HEAVY], {'in_flight': 2, 'limit': 4, 'bulk_limit': 3}) + self.assertEqual(state[LANE_ES2_KNN], {'in_flight': 1, 'limit': 3}) + self.assertEqual(state['tasks'], {'task-a': {'in_flight': 2, 'limit': 2}}) + self.assertEqual(state['tiers']['preview'], {'in_flight': 1, 'limit': 2}) + self.assertEqual(state['tiers']['core'], {'in_flight': 1, 'limit': 4}) + self.assertEqual(state['users'][str(gate.request.user.id)], {'in_flight': 1}) + + with patch.object(RedisLanes, 'find_keys', side_effect=redis.ConnectionError('down')): + self.assertEqual(get_status()['redis'], 'ConnectionError: down') + + +class CapacityViewsTest(CapacityTestMixin, OCLAPITestCase): + def setUp(self): + super().setUp() + self.staff = make_user(is_staff=True) + self.preview_user = make_user(PREVIEW_GROUP_NAME) + + def post(self, path, data, user=None, **extra): + return self.client.post( + path, data, format='json', HTTP_AUTHORIZATION=f'Token {(user or self.preview_user).get_token()}', **extra) + + def match(self, query='?semantic=true', user=None, rows=None, **extra): + return self.post( + f'/concepts/$match/{query}', + {'rows': rows or [{'name': 'a'}, {'name': 'b'}], 'target_repo_url': '/orgs/org/sources/src/'}, + user=user, **extra) + + def test_cors_exposes_the_capacity_headers(self): + for header in HEADERS: + self.assertIn(header, settings.CORS_EXPOSE_HEADERS) + + @patch('core.concepts.views.MetadataToConceptsListView.filter_queryset', return_value=[]) + def test_semantic_match_is_counted_and_carries_the_headers(self, filter_queryset_mock): + response = self.match() + + self.assertEqual(response.status_code, 200) + filter_queryset_mock.assert_called_once() + self.assertEqual(response[HEADER_DECISION], 'admitted') + self.assertEqual(response[HEADER_TIER], 'preview') + self.assertEqual(response[HEADER_IN_FLIGHT], '1') + self.assertEqual(self.redis.zcard(lane_key(LANE_API_HEAVY)), 0) # released + line = self.lines()[0] + self.assertEqual((line['endpoint'], line['rows'], line['semantic']), ('$match', 2, True)) + + @patch('core.concepts.views.MetadataToConceptsListView.filter_queryset', return_value=[]) + def test_lexical_match_is_not_a_heavy_call(self, _): + response = self.match(query='') + + self.assertEqual(response.status_code, 200) + self.assertFalse(response.has_header(HEADER_DECISION)) + self.assertEqual(self.lines(), []) + + @patch('core.concepts.views.MetadataToConceptsListView.filter_queryset', return_value=[]) + def test_reranked_lexical_match_is_a_heavy_call(self, _): + response = self.match(query='?reranker=true') + + self.assertEqual(response[HEADER_DECISION], 'admitted') + self.assertEqual(self.lines()[0]['semantic'], False) + self.assertIsNone(self.lines()[0]['es2_knn_limit']) + + @patch('core.concepts.views.MetadataToConceptsListView.filter_queryset', return_value=[]) + def test_enforce_refuses_a_match_before_charging_any_quota(self, filter_queryset_mock): + from core.capabilities.models import UsageEvent + self.configure(mode='enforce', enforce_for='all', per_user={'preview': 0}) + + response = self.match() + + self.assertEqual(response.status_code, 429) + self.assertEqual(response['Retry-After'], '120') + self.assertEqual(response[HEADER_DECISION], 'refused') + self.assertEqual(response.data['error_code'], 'capacity_exceeded') + filter_queryset_mock.assert_not_called() + self.assertFalse(UsageEvent.objects.filter(user=self.preview_user).exists()) + + @patch('core.concepts.views.MetadataToConceptsListView.filter_queryset', return_value=[]) + def test_shadow_mode_charges_and_matches_as_usual(self, filter_queryset_mock): + from core.capabilities.models import UsageEvent + self.configure(per_user={'preview': 0}) + + response = self.match() + + self.assertEqual(response.status_code, 200) + self.assertEqual(response[HEADER_DECISION], 'shadow-refused') + self.assertFalse(response.has_header('Retry-After')) + filter_queryset_mock.assert_called_once() + self.assertTrue(UsageEvent.objects.filter(user=self.preview_user, action='match_concepts').exists()) + + @patch('core.concepts.views.MetadataToConceptsListView.filter_queryset', side_effect=Exception('es down')) + def test_a_failed_match_still_releases_its_lease(self, _): + with self.assertRaises(Exception): + self.match() + + self.assertEqual(self.redis.zcard(lane_key(LANE_API_HEAVY)), 0) + self.assertEqual(self.lines()[0]['decision'], 'admitted') + + @patch('core.concepts.views.Reranker') + def test_rerank_is_counted_and_carries_the_headers(self, reranker_mock): + reranker_mock.return_value.rerank.return_value = [{'id': 1}] + + response = self.post('/concepts/$rerank/', {'rows': [{'id': 1}, {'id': 2}], 'q': 'text'}) + + self.assertEqual(response.status_code, 200) + self.assertEqual(response[HEADER_DECISION], 'admitted') + line = self.lines()[0] + self.assertEqual((line['endpoint'], line['rows'], line['single_row']), ('$rerank', 2, False)) + + @patch('core.concepts.views.Reranker') + def test_enforce_refuses_a_rerank_before_any_work(self, reranker_mock): + self.configure(mode='enforce', enforce_for='all', per_user={'preview': 0}) + + response = self.post('/concepts/$rerank/', {'rows': [{'id': 1}], 'q': 'text'}) + + self.assertEqual(response.status_code, 429) + reranker_mock.assert_not_called() + + def test_an_invalid_rerank_is_not_counted(self): + response = self.post('/concepts/$rerank/', {'rows': [{'id': 1}]}) + + self.assertEqual(response.status_code, 400) + self.assertFalse(response.has_header(HEADER_DECISION)) + + def test_config_is_staff_only(self): + self.assertEqual(self.client.get('/capacity/config/').status_code, 401) + for path in ['/capacity/config/', '/capacity/config/history/', '/capacity/status/']: + response = self.client.get(path, HTTP_AUTHORIZATION=f'Token {self.preview_user.get_token()}') + self.assertEqual(response.status_code, 403) + response = self.client.patch( + '/capacity/config/', {'mode': 'off'}, format='json', + HTTP_AUTHORIZATION=f'Token {self.preview_user.get_token()}') + self.assertEqual(response.status_code, 403) + self.assertFalse(CapacityConfig.objects.exists()) + + def test_get_config(self): + response = self.client.get('/capacity/config/', HTTP_AUTHORIZATION=f'Token {self.staff.get_token()}') + + self.assertEqual(response.status_code, 200) + self.assertIsNone(response.data['version']) + self.assertEqual(response.data['config'], get_defaults()) + self.assertEqual(response.data['defaults'], get_defaults()) + + def test_change_config_and_read_its_history(self): + auth = {'HTTP_AUTHORIZATION': f'Token {self.staff.get_token()}'} + + response = self.client.patch( + '/capacity/config/', {'tiers': {'preview': 1}, 'note': 'busy'}, format='json', **auth) + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.data['changes'], {'tiers.preview': [2, 1]}) + self.assertEqual((response.data['created_by'], response.data['note']), (self.staff.username, 'busy')) + self.assertEqual(get_config()['tiers']['preview'], 1) + + response = self.client.put('/capacity/config/', {'mode': 'off'}, format='json', **auth) + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.data['changes'], {'mode': ['shadow', 'off'], 'tiers.preview': [1, 2]}) + + response = self.client.get('/capacity/config/history/?limit=1', **auth) + self.assertEqual(len(response.data), 1) + self.assertEqual(response.data[0]['changes'], {'mode': ['shadow', 'off'], 'tiers.preview': [1, 2]}) + response = self.client.get('/capacity/config/history/', **auth) + self.assertEqual([change['changes'] for change in response.data][1], {'tiers.preview': [2, 1]}) + + def test_change_config_rejects_bad_input(self): + auth = {'HTTP_AUTHORIZATION': f'Token {self.staff.get_token()}'} + + response = self.client.patch('/capacity/config/', {'es2_knn': -1}, format='json', **auth) + self.assertEqual(response.status_code, 400) + self.assertIn('"es2_knn" must be a whole number from 0 to 100000.', response.data['detail']) + + response = self.client.patch('/capacity/config/', {'note': 5}, format='json', **auth) + self.assertEqual(response.status_code, 400) + + response = self.client.patch('/capacity/config/', [1], format='json', **auth) + self.assertEqual(response.status_code, 400) + self.assertFalse(CapacityConfig.objects.exists()) + + def test_status_view(self): + self.acquire(self.preview_user) + + response = self.client.get('/capacity/status/', HTTP_AUTHORIZATION=f'Token {self.staff.get_token()}') + + self.assertEqual(response.status_code, 200) + self.assertEqual(response.data[LANE_API_HEAVY]['in_flight'], 1) + self.assertEqual(response.data['mode'], 'shadow') + + +class CapacityCommandTest(CapacityTestMixin, OCLTestCase): + def run_command(self, *args): + out = StringIO() + call_command('capacity', *args, stdout=out) + return out.getvalue() + + def test_show_set_reset_and_history(self): + self.assertEqual(json.loads(self.run_command('show'))['config'], get_defaults()) + + output = json.loads(self.run_command('set', 'tiers.preview=1', 'mode=enforce', '--note', 'test')) + + self.assertEqual(output['changes'], {'mode': ['shadow', 'enforce'], 'tiers.preview': [2, 1]}) + latest = CapacityConfig.get_latest() + self.assertEqual((latest.source, latest.note, latest.created_by), ('command', 'test', None)) + + user = make_user(is_staff=True) + output = json.loads(self.run_command('reset', '--user', user.username)) + self.assertEqual(output['changes'], {'mode': ['enforce', 'shadow'], 'tiers.preview': [1, 2]}) + self.assertEqual(CapacityConfig.get_latest().created_by, user) + + history = self.run_command('history', '--limit', '5') + self.assertIn('"tiers.preview"', history) + self.assertIn(user.username, history) + + def test_status(self): + self.acquire(make_user(PREVIEW_GROUP_NAME)) + + self.assertEqual(json.loads(self.run_command('status'))[LANE_API_HEAVY]['in_flight'], 1) + + def test_errors(self): + for args, message in [ + (['set'], 'Give at least one'), + (['set', 'mode'], 'Expected ='), + (['set', '=1'], 'Expected ='), + (['set', 'es2_knn=-1'], '"es2_knn" must be a whole number'), + (['set', 'mode=off', '--user', 'nobody'], 'No user "nobody".'), + ]: + with self.subTest(args=args): + with self.assertRaises(CommandError) as context: + self.run_command(*args) + self.assertIn(message, str(context.exception)) + self.assertFalse(CapacityConfig.objects.exists()) diff --git a/core/capacity/urls.py b/core/capacity/urls.py new file mode 100644 index 000000000..7c8addca0 --- /dev/null +++ b/core/capacity/urls.py @@ -0,0 +1,9 @@ +from django.urls import path + +from core.capacity.views import CapacityConfigView, CapacityConfigHistoryView, CapacityStatusView + +urlpatterns = [ + path('config/', CapacityConfigView.as_view(), name='capacity-config'), + path('config/history/', CapacityConfigHistoryView.as_view(), name='capacity-config-history'), + path('status/', CapacityStatusView.as_view(), name='capacity-status'), +] diff --git a/core/capacity/views.py b/core/capacity/views.py new file mode 100644 index 000000000..c09d9bb99 --- /dev/null +++ b/core/capacity/views.py @@ -0,0 +1,93 @@ +from django.core.exceptions import ValidationError +from drf_yasg import openapi +from drf_yasg.utils import swagger_auto_schema +from rest_framework import status +from rest_framework.permissions import IsAdminUser +from rest_framework.response import Response +from rest_framework.views import APIView + +from core.capacity.config import get_current, get_defaults, save_config, diff +from core.capacity.constants import SOURCE_API +from core.capacity.limiter import get_status +from core.capacity.models import CapacityConfig +from core.common.utils import to_int + +CONFIG_BODY = openapi.Schema( + type=openapi.TYPE_OBJECT, + description='Capacity-limit settings, e.g. {"mode": "shadow", "tiers": {"preview": 1}}, plus an optional "note" ' + 'recorded with the change.' +) + + +def version_data(row): + if row is None: + return {'version': None, 'created_by': None, 'created_at': None, 'source': None, 'note': None} + return { + 'version': row.id, + 'created_by': row.created_by.username if row.created_by else None, + 'created_at': row.created_at, + 'source': row.source, + 'note': row.note or None, + } + + +class CapacityConfigView(APIView): + """ + The capacity limit on heavy calls (semantic $match, $rerank), for staff. GET returns the config in force, PATCH + merges changes into it, and PUT replaces it (anything left out takes its default). A change applies across the + API within CAPACITY_CONFIG_CACHE_SECONDS, with no deploy or restart, and is kept in the history. + """ + permission_classes = (IsAdminUser,) + + @swagger_auto_schema(operation_summary='The capacity-limit config in force (staff)') + def get(self, _): + row, config = get_current() + return Response({**version_data(row), 'config': config, 'defaults': get_defaults()}) + + @swagger_auto_schema(operation_summary='Change the capacity-limit config (staff)', request_body=CONFIG_BODY) + def patch(self, request): + return self.save(request, replace=False) + + @swagger_auto_schema(operation_summary='Replace the capacity-limit config (staff)', request_body=CONFIG_BODY) + def put(self, request): + return self.save(request, replace=True) + + @staticmethod + def save(request, replace): + if not isinstance(request.data, dict): + return Response({'detail': 'Send a JSON object.'}, status=status.HTTP_400_BAD_REQUEST) + changes = dict(request.data) + note = changes.pop('note', '') + if not isinstance(note, str): + return Response({'detail': '"note" must be a string.'}, status=status.HTTP_400_BAD_REQUEST) + try: + row, config, changed = save_config( + changes, user=request.user, source=SOURCE_API, note=note, replace=replace) + except ValidationError as ex: + return Response({'detail': ex.messages}, status=status.HTTP_400_BAD_REQUEST) + return Response({**version_data(row), 'config': config, 'changes': changed}) + + +class CapacityConfigHistoryView(APIView): + """Every change to the capacity-limit config, newest first: who, when, and each value's old and new setting.""" + permission_classes = (IsAdminUser,) + + @swagger_auto_schema( + operation_summary='Changes to the capacity-limit config (staff)', + manual_parameters=[openapi.Parameter('limit', openapi.IN_QUERY, type=openapi.TYPE_INTEGER, default=50)] + ) + def get(self, request): + limit = min(max(to_int(request.query_params.get('limit'), 50), 1), 500) + rows = CapacityConfig.objects.select_related('created_by').order_by('-id')[:limit] + return Response( + [{**version_data(row), 'changes': diff(row.previous_config, row.config)} for row in rows]) + + +class CapacityStatusView(APIView): + """Heavy calls in flight now, in each lane, against its limit (staff).""" + permission_classes = (IsAdminUser,) + + @swagger_auto_schema(operation_summary='Heavy calls in flight now (staff)') + def get(self, _): + _, config = get_current() + return Response(get_status(config)) diff --git a/core/common/utils.py b/core/common/utils.py index bc3c86e22..6a5aa4464 100644 --- a/core/common/utils.py +++ b/core/common/utils.py @@ -950,6 +950,18 @@ def parse_id(value): return None +def get_event_metadata(request): + """The request's X-OCL-Event-Metadata JSON object (oclmap's attribution bag), or {} if it's missing or malformed.""" + raw = request.META.get('HTTP_X_OCL_EVENT_METADATA') + if not raw: + return {} + try: + metadata = json.loads(raw) + except (TypeError, ValueError): + return {} + return metadata if isinstance(metadata, dict) else {} + + def generic_sort(_list): def compare(item): if isinstance(item, (int, float, str, bool)): diff --git a/core/concepts/views.py b/core/concepts/views.py index d6331577e..01c8af7e9 100644 --- a/core/concepts/views.py +++ b/core/concepts/views.py @@ -1,4 +1,3 @@ -import json import time from cid.locals import get_cid @@ -30,6 +29,8 @@ from core.capabilities.constants import CAPABILITY_EXCEEDED_ERROR_CODE, CAPABILITY_NOT_ENTITLED_ERROR_CODE, \ MAPPER_MATCH_OPERATIONS_CAPABILITY, MAPPER_MATCH_OPERATIONS_CAPABILITY_ID from core.capabilities.exceptions import CapabilityExceeded +from core.capacity.constants import ENDPOINT_MATCH, ENDPOINT_RERANK +from core.capacity.limiter import CapacityLimitMixin from core.common.permissions import CanUseMapper from core.common.search import CustomESSearch, Reranker, get_visible_repo_criteria from core.common.swagger_parameters import ( @@ -44,7 +45,7 @@ from core.common.tasks import delete_concept, make_hierarchy from core.common.throttling import ThrottleUtil from core.common.utils import (to_parent_uri_from_kwargs, generate_temp_version, get_truthy_values, to_int, - drop_version, get_falsy_values, parse_id) + drop_version, get_falsy_values, parse_id, get_event_metadata) from core.common.views import SourceChildCommonBaseView, SourceChildExtrasView, \ SourceChildExtraRetrieveUpdateDestroyView, BaseAPIView from core.concepts.constants import ( @@ -913,14 +914,8 @@ def get_match_operations_attribution(request): effort: a missing/malformed header never blocks the match, it only means the UsageEvent stays algorithm=None/map_project=None. """ - raw = request.META.get('HTTP_X_OCL_EVENT_METADATA') - if not raw: - return None, None - try: - metadata = json.loads(raw) - except (TypeError, ValueError): - return None, None - if not isinstance(metadata, dict): + metadata = get_event_metadata(request) + if not metadata: return None, None algorithm_id = metadata.get('algorithm_id') @@ -940,7 +935,7 @@ def get_match_operations_attribution(request): return algorithm_id, map_project -class MetadataToConceptsListView(BaseAPIView): # pragma: no cover +class MetadataToConceptsListView(CapacityLimitMixin, BaseAPIView): # pragma: no cover default_limit = 1 score_threshold = 0.9 score_threshold_semantic_very_high = 0.9 @@ -1208,6 +1203,18 @@ def get_repo_params(is_semantic, target_repo_params, target_repo_url, user=None) } ) def post(self, request, **kwargs): # pylint: disable=unused-argument + # Semantic (kNN) and reranked matches are heavy calls: the capacity limit counts them, and in enforce mode + # refuses them with a 429 before any quota is charged (OpenConceptLab/ocl_online#275). + rows = request.data.get('rows') + semantic = request.query_params.get('semantic', None) in TRUTHY + reranker = request.query_params.get('reranker', None) in TRUTHY + if not (isinstance(rows, list) and rows and (semantic or reranker)): + return self.match(request) + with self.capacity_gate( + request, endpoint=ENDPOINT_MATCH, rows=len(rows), semantic=semantic, reranker=reranker) as gate: + return gate.get_refusal_response() if gate.refused else self.match(request) + + def match(self, request): rows = request.data.get('rows') consumed_units = 0 if isinstance(rows, list) and rows: @@ -1252,7 +1259,7 @@ def post(self, request, **kwargs): # pylint: disable=unused-argument return response -class RerankConceptsListView(BaseAPIView): +class RerankConceptsListView(CapacityLimitMixin, BaseAPIView): is_searchable = False serializer_class = ConceptListSerializer permission_classes = (IsAuthenticated, CanUseMapper) @@ -1280,14 +1287,18 @@ def post(self, request, **kwargs): # pylint: disable=unused-argument,too-many-r {'detail': 'Missing "q" in request body.'}, status=status.HTTP_400_BAD_REQUEST ) - try: - reranker = Reranker(model_name=encoder_model) - results = reranker.rerank(hits=rows, name_key=name_key, txt=text, score_key=score_key, order_results=True) - return Response(results) - except (ValueError, RuntimeError, OSError) as e: - return Response({'detail': str(e)}, status=status.HTTP_400_BAD_REQUEST) - except Exception as e: - ERRBIT_LOGGER.log(e) - return Response( - {'detail': 'An error occurred while processing the rerank request.'}, - status=status.HTTP_500_INTERNAL_SERVER_ERROR) + with self.capacity_gate(request, endpoint=ENDPOINT_RERANK, rows=len(rows)) as gate: + if gate.refused: + return gate.get_refusal_response() + try: + reranker = Reranker(model_name=encoder_model) + results = reranker.rerank( + hits=rows, name_key=name_key, txt=text, score_key=score_key, order_results=True) + return Response(results) + except (ValueError, RuntimeError, OSError) as e: + return Response({'detail': str(e)}, status=status.HTTP_400_BAD_REQUEST) + except Exception as e: + ERRBIT_LOGGER.log(e) + return Response( + {'detail': 'An error occurred while processing the rerank request.'}, + status=status.HTTP_500_INTERNAL_SERVER_ERROR) diff --git a/core/settings.py b/core/settings.py index b81cbe956..273edab25 100644 --- a/core/settings.py +++ b/core/settings.py @@ -96,6 +96,13 @@ def get_set_from_env(name): 'X-LimitRemaining-Minute', 'X-LimitRemaining-Day', 'Retry-After', + 'X-OCL-Capacity-Decision', + 'X-OCL-Capacity-Limit', + 'X-OCL-Capacity-In-Flight', + 'X-OCL-Capacity-Tier', + 'X-OCL-Capacity-Tier-Limit', + 'X-OCL-Capacity-Tier-In-Flight', + 'X-OCL-Capacity-Suggested-Concurrency', ) CORS_ORIGIN_ALLOW_ALL = True @@ -139,6 +146,7 @@ def get_set_from_env(name): 'core.events', 'core.map_projects', 'core.capabilities', + 'core.capacity', 'core.graphql.apps.GraphqlConfig' ] REST_FRAMEWORK = { @@ -671,6 +679,14 @@ def get_set_from_env(name): ENCODER = CrossEncoder(ENCODER_MODEL_NAME, device="cpu", max_length=128) +# Capacity limit on heavy calls (semantic $match, $rerank): OpenConceptLab/ocl_online#275. The mode and numbers +# are runtime config (core.capacity.config); CAPACITY_LIMIT_MODE is only the default until staff change it. +CAPACITY_LIMIT_MODE = os.environ.get('CAPACITY_LIMIT_MODE', 'shadow') +CAPACITY_CONFIG_CACHE_SECONDS = int(os.environ.get('CAPACITY_CONFIG_CACHE_SECONDS', 10)) +CAPACITY_REDIS_TIMEOUT_SECONDS = float(os.environ.get('CAPACITY_REDIS_TIMEOUT_SECONDS', 0.5)) +CAPACITY_REDIS_RETRY_SECONDS = int(os.environ.get('CAPACITY_REDIS_RETRY_SECONDS', 30)) +CAPACITY_TASK_ID = os.environ.get('CAPACITY_TASK_ID', '') + ANALYTICS_API = os.environ.get('ANALYTICS_API', 'http://host.docker.internal:8002') if ANALYTICS_API: MIDDLEWARE = [*MIDDLEWARE, 'core.middlewares.middlewares.AnalyticsMiddleware'] diff --git a/core/urls.py b/core/urls.py index 415ff2ed1..fbb1df547 100644 --- a/core/urls.py +++ b/core/urls.py @@ -118,6 +118,7 @@ path('manage/bulkimport/', BulkImportView.as_view(), name='bulk_import_urls'), path('toggles/', include('core.toggles.urls'), name='toggles'), path('capabilities/', include('core.capabilities.urls'), name='capabilities'), + path('capacity/', include('core.capacity.urls'), name='capacity'), ] if ENV == 'development': diff --git a/core/users/constants.py b/core/users/constants.py index acd96f479..fbcbbd498 100644 --- a/core/users/constants.py +++ b/core/users/constants.py @@ -28,6 +28,7 @@ MAPPER_SCISPACY_PERMISSION = 'users.mapper_scispacy' PREVIEW_GROUP = 'preview' PREVIEW_GRANDFATHERED_GROUP = 'preview_grandfathered' # existing accounts (ocl_online#230); layered on `preview` +EARLY_ACCESS_GROUP = 'early_access' BULK_IMPORT_ADVANCED_PERMISSION = 'users.bulk_import_advanced' BULK_IMPORT_PRIORITY_PERMISSION = 'users.bulk_import_priority' LIST_UNPAGINATED_PERMISSION = 'users.list_unpaginated' diff --git a/requirements.txt b/requirements.txt index 0e9cd3dae..75e411705 100644 --- a/requirements.txt +++ b/requirements.txt @@ -30,6 +30,8 @@ django-ordered-model==3.7.4 django-health-check==3.17.0 markdown==3.8.1 mock==5.1.0 +fakeredis[lua]==2.38.0 # tests run the capacity limit's Redis scripts without a Redis server +lupa==2.8 django-request-logging==0.7.5 django-cid==2.4 django-dirtyfields==1.9.2 From 626b66693758bbb92f843613476b5da13cff3110 Mon Sep 17 00:00:00 2001 From: Jonathan Payne Date: Tue, 29 Sep 2026 23:22:18 -0400 Subject: [PATCH 2/8] OpenConceptLab/ocl_online#275 | Capacity limit: fixes from Codex review, pass 1 - A lane's TTL is only ever extended, so a shorter lease (after a runtime change) can't expire a lane that holds longer leases. - If the renewer thread fails to start, the leases are still released: stopping the renewer and releasing are separate steps. - Renewal and release skip Redis while it's marked failed (the leases expire), and release no longer waits on a renewal in flight, which can only extend leases the call still holds. - A config that can't be re-read never refuses: a cached enforce config falls back to shadow until a read succeeds. - A config change that's saved no longer fails the request if its log line can't be written. - renew_seconds must be at most a third of lease_seconds, and lease_seconds at least 15. - Sentinel connections are explicitly retry-free; the trace header parser tolerates whitespace. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TyAC9YAn5hP3yGSREkSN8m --- core/capacity/config.py | 28 ++++++++------ core/capacity/limiter.py | 52 +++++++++++++++++--------- core/capacity/tests/tests.py | 71 ++++++++++++++++++++++++++++++++---- 3 files changed, 114 insertions(+), 37 deletions(-) diff --git a/core/capacity/config.py b/core/capacity/config.py index 58181530b..574a1fbcf 100644 --- a/core/capacity/config.py +++ b/core/capacity/config.py @@ -14,7 +14,7 @@ from django.db import connection, transaction from core.capacity.constants import ( - MODES, MODE_SHADOW, ENFORCE_FOR, ENFORCE_FOR_AWARE, TIER_STAFF, TIER_CORE, TIER_EARLY_ACCESS, + MODES, MODE_SHADOW, MODE_ENFORCE, ENFORCE_FOR, ENFORCE_FOR_AWARE, TIER_STAFF, TIER_CORE, TIER_EARLY_ACCESS, TIER_PREVIEW, CONFIG_LOG_EVENT, SOURCE_API) from core.capacity.logs import emit @@ -111,10 +111,10 @@ def validate(config): if not errors: if config['reserve_single_row'] > config['api_heavy']['cluster']: errors.append('"reserve_single_row" can\'t be more than "api_heavy.cluster".') - if config['lease_seconds'] < 5: - errors.append('"lease_seconds" must be at least 5.') - if not 1 <= config['renew_seconds'] < config['lease_seconds']: - errors.append('"renew_seconds" must be at least 1 and less than "lease_seconds".') + if config['lease_seconds'] < 15: + errors.append('"lease_seconds" must be at least 15.') + if not 1 <= config['renew_seconds'] <= config['lease_seconds'] // 3: + errors.append('"renew_seconds" must be at least 1 and at most a third of "lease_seconds".') if config['max_hold_seconds'] < config['lease_seconds']: errors.append('"max_hold_seconds" can\'t be less than "lease_seconds".') retry_after = config['retry_after'] @@ -154,9 +154,12 @@ def get_config(): try: _, config = get_current() except Exception as ex: - # Keep the last good config (or the defaults) rather than fail the request. + # Keep the last good config (or the defaults) rather than fail the request, but never refuse on a + # config that couldn't be confirmed: enforce falls back to shadow until a read succeeds. logger.warning('Capacity config could not be read (%s); using the last known config', ex) - config = _cache['config'] or get_defaults() + config = copy.deepcopy(_cache['config'] or get_defaults()) + if config['mode'] == MODE_ENFORCE: + config['mode'] = MODE_SHADOW _cache['config'] = config _cache['expires_at'] = now + settings.CAPACITY_CONFIG_CACHE_SECONDS return config @@ -187,8 +190,11 @@ def save_config(changes, user=None, source=SOURCE_API, note='', replace=False): row = CapacityConfig.objects.create( config=config, previous_config=previous, created_by=user, source=source, note=note or '') clear_cache() - emit({ - 'event': CONFIG_LOG_EVENT, 'version': row.id, 'changed_by': getattr(user, 'username', None), - 'source': source, 'note': note or None, 'changes': changed, 'mode': config['mode'], - }) + try: + emit({ + 'event': CONFIG_LOG_EVENT, 'version': row.id, 'changed_by': getattr(user, 'username', None), + 'source': source, 'note': note or None, 'changes': changed, 'mode': config['mode'], + }) + except Exception as ex: # the change is saved, and its row is the record; don't fail the request over a log line + logger.warning('Capacity config version %s was saved, but its log line failed: %s', row.id, ex) return row, config, changed diff --git a/core/capacity/limiter.py b/core/capacity/limiter.py index d0332f02b..6937ecdc6 100644 --- a/core/capacity/limiter.py +++ b/core/capacity/limiter.py @@ -1,7 +1,8 @@ """ The capacity limit on heavy calls (OpenConceptLab/ocl_online#275): semantic `$match` (kNN searches and/or the in-request rerank) and `$rerank`. It caps how many run at once across all users, to protect the API workers and -Elasticsearch. It isn't a quota: it charges nothing, and a call it refuses (in enforce mode only) is asked to retry shortly. +Elasticsearch. It isn't a quota: it charges nothing, and a call it refuses (in enforce mode only) is asked to retry +shortly. Each lane is a Redis sorted set of leases: member = a random token per call, score = the lease's expiry in ms by Redis's own clock, so the API tasks' clocks don't matter. A call takes a lease in every lane that applies to it, in @@ -9,8 +10,9 @@ that isn't renewed expires, so a worker that dies frees its slots within `lease_seconds`. If Redis can't be reached, calls go ahead uncounted and it's logged (fail open): the limiter must never cause an -outage. It uses its own Redis client with short timeouts and no retries, and skips Redis for -CAPACITY_REDIS_RETRY_SECONDS after an error, so an outage costs a call at most one short timeout. +outage. It uses its own Redis client with short timeouts and no retries, and each process skips Redis (acquire, +renewal and release) for CAPACITY_REDIS_RETRY_SECONDS after an error. So an outage delays at most one call per +process in that time, by a timeout per connection attempt (a few, with Sentinel); leases it can't release expire. """ import socket import threading @@ -42,6 +44,10 @@ if redis.replicate_commands then redis.replicate_commands() end local now = redis.call('TIME') local now_ms = tonumber(now[1]) * 1000 + math.floor(tonumber(now[2]) / 1000) +-- Extend a lane's TTL to cover a new lease, never shorten it: other calls' leases may be longer. +local function keep(key, ms) + if redis.call('PTTL', key) < ms then redis.call('PEXPIRE', key, ms) end +end """ # KEYS: one per lane. ARGV: token, lease ms, force ('1': take the lease even when a lane is full), then one limit @@ -60,7 +66,7 @@ result[1] = 1 for _, key in ipairs(KEYS) do redis.call('ZADD', key, now_ms + lease_ms, ARGV[1]) - redis.call('PEXPIRE', key, lease_ms * 2) + keep(key, lease_ms * 2) end end return result @@ -73,7 +79,7 @@ for _, key in ipairs(KEYS) do if redis.call('ZSCORE', key, ARGV[1]) then redis.call('ZADD', key, now_ms + lease_ms, ARGV[1]) - redis.call('PEXPIRE', key, lease_ms * 2) + keep(key, lease_ms * 2) renewed = renewed + 1 end end @@ -106,7 +112,8 @@ def build_redis_client(): 'retry_on_timeout': False, 'health_check_interval': 0, } if settings.REDIS_SENTINELS: - sentinel_options = {'socket_timeout': timeout, 'socket_connect_timeout': timeout} + sentinel_options = { + 'socket_timeout': timeout, 'socket_connect_timeout': timeout, 'retry': Retry(NoBackoff(), 0)} if settings.REDIS_PASSWORD: sentinel_options['password'] = settings.REDIS_PASSWORD sentinel = Sentinel(settings.REDIS_SENTINELS_LIST, sentinel_kwargs=sentinel_options) @@ -201,8 +208,8 @@ def get_queue_ms(request, now=None): adds X-Amzn-Trace-Id with the time it received the request, in whole seconds (in `Self`, or in `Root` when it started the trace), so this reads up to a second high. Mostly it's time spent queued for a free API worker. """ - fields = dict(part.split('=', 1) for part in (request.META.get('HTTP_X_AMZN_TRACE_ID') or '').split(';') - if '=' in part) + fields = {key.strip(): value.strip() for key, value in ( + part.split('=', 1) for part in (request.META.get('HTTP_X_AMZN_TRACE_ID') or '').split(';') if '=' in part)} try: received = int((fields.get('Self') or fields.get('Root')).split('-')[1], 16) except (AttributeError, IndexError, ValueError): @@ -225,8 +232,10 @@ def run(self): self.gate.renew() def stop(self): + # No need to wait out a renewal in flight: renewing only extends leases the call still holds, so one that + # lands after the release changes nothing. self.finished.set() - self.join(timeout=2 * settings.CAPACITY_REDIS_TIMEOUT_SECONDS + 1) + self.join(timeout=0.2) class CapacityGate: @@ -303,8 +312,9 @@ def acquire(self): self.decision = DECISION_SHADOW_REFUSED if taken else DECISION_REFUSED self.retry_after = self.get_retry_after(dict(zip([lane for lane, _, _ in self.lanes], before))) if taken: - self.renewer = LeaseRenewer(self) - self.renewer.start() + renewer = LeaseRenewer(self) + renewer.start() + self.renewer = renewer except Exception as ex: self.record_error(ex) self.decision = self.decision or DECISION_UNAVAILABLE @@ -353,6 +363,8 @@ def get_suggested_concurrency(self): def renew(self): try: + if not RedisLanes.is_available(): + return if RedisLanes.renew(self.keys, self.token, self.lease_ms) < len(self.lanes): self.lease_lost = True except Exception as ex: @@ -365,16 +377,20 @@ def release(self): try: if self.renewer: self.renewer.stop() - if self.holding: + except Exception as ex: + self.record_error(ex) + try: + if self.holding and not RedisLanes.is_available(): + self.error = self.error or 'release skipped: Redis failed recently; the leases will expire' + elif self.holding: RedisLanes.release(self.keys, self.token) except Exception as ex: self.record_error(ex) - finally: - if self.decision: - try: - self.log() - except Exception: # a log line must never fail the call - pass + if self.decision: + try: + self.log() + except Exception: # a log line must never fail the call + pass def record_error(self, ex): if isinstance(ex, redis.RedisError): diff --git a/core/capacity/tests/tests.py b/core/capacity/tests/tests.py index 550c82bc4..ae938b9e6 100644 --- a/core/capacity/tests/tests.py +++ b/core/capacity/tests/tests.py @@ -114,9 +114,9 @@ def test_validate_rejects_bad_configs(self): ({'es2_knn': '3'}, '"es2_knn" must be a whole number'), ({'per_user': {'preview': 100001}}, '"per_user.preview" must be a whole number'), ({'reserve_single_row': 5}, '"reserve_single_row" can\'t be more than "api_heavy.cluster".'), - ({'lease_seconds': 4, 'renew_seconds': 1}, '"lease_seconds" must be at least 5.'), - ({'renew_seconds': 60}, '"renew_seconds" must be at least 1 and less than "lease_seconds".'), - ({'renew_seconds': 0}, '"renew_seconds" must be at least 1 and less than "lease_seconds".'), + ({'lease_seconds': 14, 'renew_seconds': 1}, '"lease_seconds" must be at least 15.'), + ({'renew_seconds': 21}, '"renew_seconds" must be at least 1 and at most a third of "lease_seconds".'), + ({'renew_seconds': 0}, '"renew_seconds" must be at least 1 and at most a third of "lease_seconds".'), ({'max_hold_seconds': 30}, '"max_hold_seconds" can\'t be less than "lease_seconds".'), ({'retry_after': {'base': 0}}, '"retry_after" needs'), ({'retry_after': {'base': 10, 'max': 5}}, '"retry_after" needs'), @@ -185,6 +185,15 @@ def test_save_config_replace_starts_from_the_defaults(self): self.assertEqual(config, {**get_defaults(), 'mode': 'off'}) self.assertEqual(changes, {'es2_knn': [1, 3], 'mode': ['shadow', 'off'], 'tiers.preview': [1, 2]}) + def test_a_failed_log_line_doesnt_fail_a_saved_change(self): + self.emit.side_effect = BrokenPipeError('stdout closed') + + with self.assertLogs('oclapi', level='WARNING'): + row, config, _ = save_config({'es2_knn': 1}) + + self.assertEqual(CapacityConfig.get_latest(), row) + self.assertEqual(config['es2_knn'], 1) + def test_save_config_rejects_an_invalid_change(self): with self.assertRaises(ValidationError): save_config({'es2_knn': -1}) @@ -205,12 +214,16 @@ def test_get_config_keeps_the_last_known_config_when_the_database_fails(self): with self.assertLogs('oclapi', level='WARNING'): self.assertEqual(get_config(), get_defaults()) - save_config({'es2_knn': 1}) + save_config({'es2_knn': 1, 'mode': 'enforce'}) get_config() _cache['expires_at'] = 0.0 with patch('core.capacity.config.get_current', side_effect=Exception('db down')): with self.assertLogs('oclapi', level='WARNING'): - self.assertEqual(get_config()['es2_knn'], 1) + config = get_config() + self.assertEqual(config['es2_knn'], 1) + self.assertEqual(config['mode'], 'shadow') # never refuse on a config that couldn't be confirmed + clear_cache() + self.assertEqual(get_config()['mode'], 'enforce') class CapacityHelpersTest(OCLTestCase): @@ -240,6 +253,8 @@ def request(trace_id=None): self.assertIsNone(get_queue_ms(request('Root=1-zz-abc'), now=received)) self.assertIsNone(get_queue_ms(request('Root=nope'), now=received)) self.assertIsNone(get_queue_ms(request('garbage'), now=received)) + self.assertEqual( + get_queue_ms(request(f'Root=1-{received - 100:x}-abc; Self=1-{received:x}-def'), now=received + 1), 1000) def test_get_event_metadata(self): def request(value=None): @@ -278,9 +293,13 @@ def test_build_redis_client_with_sentinels(self): client = build_redis_client() self.assertEqual(client, sentinel_mock.return_value.master_for.return_value) - sentinel_mock.assert_called_once_with( - [('s1', 26379), ('s2', 26379)], - sentinel_kwargs={'socket_timeout': 0.25, 'socket_connect_timeout': 0.25, 'password': 'secret'}) + args, kwargs = sentinel_mock.call_args + self.assertEqual(args, ([('s1', 26379), ('s2', 26379)],)) + sentinel_options = kwargs['sentinel_kwargs'] + self.assertEqual(sentinel_options['retry']._retries, 0) # pylint: disable=protected-access + self.assertEqual( + {key: value for key, value in sentinel_options.items() if key != 'retry'}, + {'socket_timeout': 0.25, 'socket_connect_timeout': 0.25, 'password': 'secret'}) args, kwargs = sentinel_mock.return_value.master_for.call_args self.assertEqual(args, ('primary',)) self.assertEqual((kwargs['password'], kwargs['socket_timeout']), ('secret', 0.25)) @@ -520,12 +539,48 @@ def test_a_failed_release_or_renewal_is_logged_not_raised(self): with patch.object(RedisLanes, 'renew', side_effect=redis.TimeoutError('slow')): gate.renew() self.assertEqual(gate.error, 'TimeoutError: slow') + self.assertFalse(RedisLanes.is_available()) + RedisLanes.unavailable_until = 0.0 with patch.object(RedisLanes, 'release', side_effect=redis.ConnectionError('gone')): gate.release() self.assertEqual(self.lines()[-1]['error'], 'ConnectionError: gone') self.assertFalse(RedisLanes.is_available()) + def test_renewal_and_release_skip_redis_after_it_failed(self): + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + RedisLanes.mark_unavailable() + + with patch.object(RedisLanes, 'renew') as renew_mock, patch.object(RedisLanes, 'release') as release_mock: + gate.renew() + gate.release() + + renew_mock.assert_not_called() + release_mock.assert_not_called() + self.assertEqual(self.lines()[-1]['error'], 'release skipped: Redis failed recently; the leases will expire') + self.assertEqual(self.redis.zcard(lane_key(LANE_USER, gate.request.user.id)), 1) + + def test_a_renewer_that_fails_to_start_still_releases(self): + with patch.object(LeaseRenewer, 'start', side_effect=RuntimeError("can't start new thread")): + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + + self.assertEqual(gate.decision, DECISION_ADMITTED) + self.assertIsNone(gate.renewer) + self.assertEqual(gate.error, "RuntimeError: can't start new thread") + gate.release() + for key in gate.keys: + self.assertEqual(self.redis.zcard(key), 0) + + def test_a_short_lease_never_shortens_a_lanes_ttl(self): + self.redis.zadd(lane_key(LANE_API_HEAVY), {'long-lease': time.time() * 1000 + 600000}) + self.redis.pexpire(lane_key(LANE_API_HEAVY), 1200000) + + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + gate.renew() + + self.assertGreater(self.redis.pttl(lane_key(LANE_API_HEAVY)), 600000) + self.assertGreater(self.redis.pttl(lane_key(LANE_USER, gate.request.user.id)), 60000) + def test_an_unexpected_error_fails_open(self): with patch('core.capacity.limiter.get_tier', side_effect=KeyError('tier')): gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) From 570833aff75ef79e5c5a1314b1ba349fa4abf0a0 Mon Sep 17 00:00:00 2001 From: Jonathan Payne Date: Tue, 29 Sep 2026 23:35:33 -0400 Subject: [PATCH 3/8] OpenConceptLab/ocl_online#275 | Capacity limit: fixes from Codex review, pass 2 - Every Redis operation runs on a per-process Redis thread with an overall deadline (CAPACITY_REDIS_DEADLINE_SECONDS, 1 s), which also bounds DNS and a walk through the Sentinels. An operation still queued when its caller gives up is skipped. - Log lines go through a bounded queue and a writer thread, so a backed-up log sink can't hold a call; when the queue is full a line is dropped, and the next one says how many were. - Renewal is all or nothing: leases are renewed only if the call still holds every one unexpired, so an expired lease is never revived; once a lease is lost the renewer stops. - An error inside acquire never refuses a call: a call that doesn't hold a lease becomes unavailable. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TyAC9YAn5hP3yGSREkSN8m --- core/capacity/limiter.py | 89 +++++++++++++++++----- core/capacity/logs.py | 60 ++++++++++++++- core/capacity/tests/tests.py | 144 +++++++++++++++++++++++++++++++++-- core/settings.py | 1 + 4 files changed, 267 insertions(+), 27 deletions(-) diff --git a/core/capacity/limiter.py b/core/capacity/limiter.py index 6937ecdc6..651d78ef7 100644 --- a/core/capacity/limiter.py +++ b/core/capacity/limiter.py @@ -10,14 +10,17 @@ that isn't renewed expires, so a worker that dies frees its slots within `lease_seconds`. If Redis can't be reached, calls go ahead uncounted and it's logged (fail open): the limiter must never cause an -outage. It uses its own Redis client with short timeouts and no retries, and each process skips Redis (acquire, -renewal and release) for CAPACITY_REDIS_RETRY_SECONDS after an error. So an outage delays at most one call per -process in that time, by a timeout per connection attempt (a few, with Sentinel); leases it can't release expire. +outage. It uses its own Redis client with short timeouts and no retries, and waits at most +CAPACITY_REDIS_DEADLINE_SECONDS for any Redis operation, DNS and Sentinel discovery included. After an error, each +process skips Redis (acquire, renewal and release) for CAPACITY_REDIS_RETRY_SECONDS, and leases it can't release +expire. So an outage delays at most one call per process in that time, by at most the deadline. """ +import os import socket import threading import time import uuid +from concurrent.futures import ThreadPoolExecutor, TimeoutError as FutureTimeoutError from contextlib import contextmanager import redis @@ -72,18 +75,23 @@ return result """ -# KEYS: the call's lanes. ARGV: token, lease ms. Extends the leases the call still holds; returns how many. +# KEYS: the call's lanes. ARGV: token, lease ms. Renews the call's leases only if it still holds every one of them +# unexpired, so a lease that expired is never revived and a call is never counted in only some lanes. Returns how +# many it still held. RENEW_SCRIPT = _NOW_MS + """ local lease_ms = tonumber(ARGV[2]) -local renewed = 0 +local held = 0 for _, key in ipairs(KEYS) do - if redis.call('ZSCORE', key, ARGV[1]) then + local expiry = redis.call('ZSCORE', key, ARGV[1]) + if expiry and tonumber(expiry) > now_ms then held = held + 1 end +end +if held == #KEYS then + for _, key in ipairs(KEYS) do redis.call('ZADD', key, now_ms + lease_ms, ARGV[1]) keep(key, lease_ms * 2) - renewed = renewed + 1 end end -return renewed +return held """ # KEYS: lanes. Returns each lane's count of unexpired leases. @@ -130,6 +138,8 @@ class RedisLanes: scripts = {} unavailable_until = 0.0 lock = threading.Lock() + executor = None + executor_pid = None @classmethod def use_client(cls, client): @@ -137,6 +147,39 @@ def use_client(cls, client): cls.client = client cls.scripts = {} cls.unavailable_until = 0.0 + if cls.executor is not None: + cls.executor.shutdown(wait=False, cancel_futures=True) + cls.executor = None + + @classmethod + def get_executor(cls): + if cls.executor is None or cls.executor_pid != os.getpid(): + with cls.lock: + if cls.executor is None or cls.executor_pid != os.getpid(): + cls.executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix='ocl-capacity-redis') + cls.executor_pid = os.getpid() + return cls.executor + + @classmethod + def call(cls, operation, *args): + """ + Run a Redis operation on this process's Redis thread, and wait at most CAPACITY_REDIS_DEADLINE_SECONDS for + it in all: socket timeouts don't bound DNS or a walk through the Sentinels. An operation still queued when + its caller gives up is skipped. One already running finishes, so an acquire that lands late holds a lease + nobody renews, which expires. + """ + abandoned = threading.Event() + + def run(): + return None if abandoned.is_set() else operation(*args) + + future = cls.get_executor().submit(run) + try: + return future.result(timeout=settings.CAPACITY_REDIS_DEADLINE_SECONDS) + except FutureTimeoutError as ex: + abandoned.set() + raise redis.TimeoutError( + f'no answer from Redis within {settings.CAPACITY_REDIS_DEADLINE_SECONDS} s') from ex @classmethod def get_client(cls): @@ -166,28 +209,33 @@ def mark_unavailable(cls): @classmethod def acquire(cls, keys, limits, token, lease_ms, force): # pylint: disable=too-many-arguments """(whether the lease was taken, each lane's count before this call)""" - result = cls.run_script(ACQUIRE_SCRIPT, keys, [token, lease_ms, '1' if force else '0', *limits]) + result = cls.call(cls.run_script, ACQUIRE_SCRIPT, keys, [token, lease_ms, '1' if force else '0', *limits]) return bool(result[0]), [int(count) for count in result[1:]] @classmethod def renew(cls, keys, token, lease_ms): - return int(cls.run_script(RENEW_SCRIPT, keys, [token, lease_ms])) + """How many of its leases the call still held; they're renewed only if that's all of them.""" + return int(cls.call(cls.run_script, RENEW_SCRIPT, keys, [token, lease_ms])) @classmethod def release(cls, keys, token): - pipeline = cls.get_client().pipeline(transaction=False) - for key in keys: - pipeline.zrem(key, token) - pipeline.execute() + def remove(): + pipeline = cls.get_client().pipeline(transaction=False) + for key in keys: + pipeline.zrem(key, token) + pipeline.execute() + cls.call(remove) @classmethod def count(cls, keys): - return [int(count) for count in cls.run_script(COUNT_SCRIPT, keys)] if keys else [] + return [int(count) for count in cls.call(cls.run_script, COUNT_SCRIPT, keys)] if keys else [] @classmethod def find_keys(cls, lane): - return sorted(key.decode() if isinstance(key, bytes) else key - for key in cls.get_client().scan_iter(match=lane_key(lane, '*'), count=100)) + def scan(): + return sorted(key.decode() if isinstance(key, bytes) else key + for key in cls.get_client().scan_iter(match=lane_key(lane, '*'), count=100)) + return cls.call(scan) def get_tier(user): @@ -230,12 +278,14 @@ def __init__(self, gate): def run(self): while not self.finished.wait(self.interval) and time.monotonic() < self.deadline: self.gate.renew() + if self.gate.lease_lost: + return # its leases expired: renewing the rest would count the call in only some lanes def stop(self): # No need to wait out a renewal in flight: renewing only extends leases the call still holds, so one that # lands after the release changes nothing. self.finished.set() - self.join(timeout=0.2) + self.join(timeout=0.1) class CapacityGate: @@ -317,7 +367,8 @@ def acquire(self): self.renewer = renewer except Exception as ex: self.record_error(ex) - self.decision = self.decision or DECISION_UNAVAILABLE + if not self.holding or not self.decision: # never refuse because the limiter itself failed + self.decision, self.retry_after = DECISION_UNAVAILABLE, None return self def get_lanes(self): diff --git a/core/capacity/logs.py b/core/capacity/logs.py index 841d32f65..54afcc949 100644 --- a/core/capacity/logs.py +++ b/core/capacity/logs.py @@ -1,10 +1,66 @@ +import atexit import json +import os +import queue +import sys +import threading + +MAX_QUEUED_LINES = 1000 + +_state = {'queue': None, 'pid': None, 'dropped': 0} +_lock = threading.Lock() def emit(record): """ Write one JSON object as one line of the API log, where CloudWatch metric filters match it. The API's other timing lines are prints too: gunicorn captures stdout, and the `oclapi` logger has no handler in production. + + Never blocks the request: the line goes on a bounded queue that a writer thread prints, so a backed-up log sink + can't hold a call. When the queue is full the line is dropped, and the next line that gets through says how + many were (`log_dropped`). """ - print(json.dumps({key: value for key, value in record.items() if value is not None}, - separators=(',', ':'), default=str), flush=True) + dropped = _state['dropped'] + if dropped: + record = {**record, 'log_dropped': dropped} + line = json.dumps({key: value for key, value in record.items() if value is not None}, + separators=(',', ':'), default=str) + try: + get_queue().put_nowait(line) + _state['dropped'] -= dropped + except queue.Full: + _state['dropped'] += 1 + + +def get_queue(): + if _state['pid'] != os.getpid(): + with _lock: + if _state['pid'] != os.getpid(): # first use in this process (gunicorn forks workers) + lines = queue.Queue(maxsize=MAX_QUEUED_LINES) + threading.Thread(target=write_lines, args=(lines,), name='ocl-capacity-log', daemon=True).start() + _state.update(queue=lines, pid=os.getpid(), dropped=0) + return _state['queue'] + + +def write_lines(lines): + while True: + write(lines.get()) + + +def write(line): + try: + sys.stdout.write(line + '\n') + sys.stdout.flush() + except Exception: + pass + + +@atexit.register +def drain(): + """Print what's still queued when a worker exits (gunicorn recycles them).""" + lines = _state['queue'] if _state['pid'] == os.getpid() else None + while lines is not None: + try: + write(lines.get_nowait()) + except queue.Empty: + break diff --git a/core/capacity/tests/tests.py b/core/capacity/tests/tests.py index ae938b9e6..ccc6c7377 100644 --- a/core/capacity/tests/tests.py +++ b/core/capacity/tests/tests.py @@ -19,6 +19,7 @@ HEADER_SUGGESTED_CONCURRENCY, ENDPOINT_MATCH, ENDPOINT_RERANK, DECISION_ADMITTED, DECISION_SHADOW_REFUSED, DECISION_REFUSED, DECISION_UNAVAILABLE, LANE_API_HEAVY, LANE_API_HEAVY_TASK, LANE_ES2_KNN, LANE_TIER, LANE_USER, LOG_EVENT, CONFIG_LOG_EVENT, SOURCE_COMMAND) +from core.capacity import logs from core.capacity.limiter import ( RedisLanes, CapacityGate, LeaseRenewer, get_tier, get_queue_ms, lane_key, get_task_id, build_redis_client, get_status) @@ -475,22 +476,42 @@ def test_suggested_concurrency(self): self.assertEqual(preview.get_suggested_concurrency(), 1) # never below 1 unless paused self.assertIsNone(CapacityGate(Mock(), ENDPOINT_MATCH).get_suggested_concurrency()) - def test_renewal_extends_the_leases_it_still_holds(self): + def test_renewal_extends_every_lease_or_none(self): gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) key = lane_key(LANE_USER, gate.request.user.id) - self.redis.zadd(key, {gate.token: 1000}) + soon = time.time() * 1000 + 5000 + for lane_key_ in gate.keys: + self.redis.zadd(lane_key_, {gate.token: soon}) gate.renew() - self.assertGreater(self.redis.zscore(key, gate.token), time.time() * 1000) + for lane_key_ in gate.keys: + self.assertGreater(self.redis.zscore(lane_key_, gate.token), soon + 30000) self.assertFalse(gate.lease_lost) - self.redis.delete(key) + + # one lease has expired (its renewal came late): nothing is renewed, and the expired one isn't revived + self.redis.zadd(key, {gate.token: 1000}) + cluster_expiry = self.redis.zscore(lane_key(LANE_API_HEAVY), gate.token) gate.renew() self.assertTrue(gate.lease_lost) + self.assertEqual(self.redis.zscore(key, gate.token), 1000) + self.assertEqual(self.redis.zscore(lane_key(LANE_API_HEAVY), gate.token), cluster_expiry) + + self.redis.delete(key) + gate.renew() self.assertFalse(self.redis.exists(key)) # a lost lease isn't re-added + def test_the_renewer_stops_once_a_lease_is_lost(self): + gate = Mock(config={'renew_seconds': 0.01, 'max_hold_seconds': 60}, lease_lost=True) + renewer = LeaseRenewer(gate) + renewer.start() + renewer.join(1) + + self.assertFalse(renewer.is_alive()) + gate.renew.assert_called_once() + def test_the_renewer_thread_renews_until_stopped(self): - gate = Mock(config={'renew_seconds': 0.01, 'max_hold_seconds': 60}) + gate = Mock(config={'renew_seconds': 0.01, 'max_hold_seconds': 60}, lease_lost=False) renewer = LeaseRenewer(gate) renewer.start() time.sleep(0.1) @@ -499,7 +520,7 @@ def test_the_renewer_thread_renews_until_stopped(self): self.assertFalse(renewer.is_alive()) self.assertGreater(gate.renew.call_count, 1) - gate = Mock(config={'renew_seconds': 0.01, 'max_hold_seconds': 0}) + gate = Mock(config={'renew_seconds': 0.01, 'max_hold_seconds': 0}, lease_lost=False) renewer = LeaseRenewer(gate) renewer.start() renewer.join(1) @@ -581,6 +602,48 @@ def test_a_short_lease_never_shortens_a_lanes_ttl(self): self.assertGreater(self.redis.pttl(lane_key(LANE_API_HEAVY)), 600000) self.assertGreater(self.redis.pttl(lane_key(LANE_USER, gate.request.user.id)), 60000) + @override_settings(CAPACITY_REDIS_DEADLINE_SECONDS=0.1) + def test_a_redis_call_that_hangs_fails_open_at_the_deadline(self): + def hang(*_): + time.sleep(0.5) + + with patch.object(RedisLanes, 'run_script', side_effect=hang): + started = time.monotonic() + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + elapsed = time.monotonic() - started + + self.assertLess(elapsed, 0.4) + self.assertEqual(gate.decision, DECISION_UNAVAILABLE) + self.assertEqual(gate.error, 'TimeoutError: no answer from Redis within 0.1 s') + self.assertFalse(RedisLanes.is_available()) + + def test_a_queued_redis_call_is_skipped_once_its_caller_gave_up(self): + release_first = __import__('threading').Event() + calls = [] + + def first(): + release_first.wait(1) + + with override_settings(CAPACITY_REDIS_DEADLINE_SECONDS=0.05): + with self.assertRaises(redis.TimeoutError): + RedisLanes.call(first) + with self.assertRaises(redis.TimeoutError): + RedisLanes.call(calls.append, 'second') # queued behind the first, which still hangs + release_first.set() + RedisLanes.call(lambda: None) # runs after both + + self.assertEqual(calls, []) + + def test_an_error_after_a_refusal_never_refuses(self): + self.configure(mode='enforce', enforce_for='all', per_user={'preview': 0}) + + with patch.object(CapacityGate, 'get_retry_after', side_effect=ValueError('bad math')): + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + + self.assertEqual(gate.decision, DECISION_UNAVAILABLE) + self.assertFalse(gate.refused) + self.assertIsNone(gate.retry_after) + def test_an_unexpected_error_fails_open(self): with patch('core.capacity.limiter.get_tier', side_effect=KeyError('tier')): gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) @@ -740,6 +803,19 @@ def test_enforce_refuses_a_match_before_charging_any_quota(self, filter_queryset filter_queryset_mock.assert_not_called() self.assertFalse(UsageEvent.objects.filter(user=self.preview_user).exists()) + @patch('core.concepts.views.MetadataToConceptsListView.filter_queryset', return_value=[]) + def test_a_limiter_error_in_enforce_mode_lets_the_match_through(self, filter_queryset_mock): + from core.capabilities.models import UsageEvent + self.configure(mode='enforce', enforce_for='all', per_user={'preview': 0}) + + with patch.object(CapacityGate, 'get_retry_after', side_effect=ValueError('bad math')): + response = self.match() + + self.assertEqual(response.status_code, 200) + self.assertEqual(response[HEADER_DECISION], 'unavailable') + filter_queryset_mock.assert_called_once() + self.assertTrue(UsageEvent.objects.filter(user=self.preview_user, action='match_concepts').exists()) + @patch('core.concepts.views.MetadataToConceptsListView.filter_queryset', return_value=[]) def test_shadow_mode_charges_and_matches_as_usual(self, filter_queryset_mock): from core.capabilities.models import UsageEvent @@ -894,3 +970,59 @@ def test_errors(self): self.run_command(*args) self.assertIn(message, str(context.exception)) self.assertFalse(CapacityConfig.objects.exists()) + + +class CapacityLogsTest(OCLTestCase): + def setUp(self): + super().setUp() + logs._state.update(queue=None, pid=None, dropped=0) # pylint: disable=protected-access + + def tearDown(self): + logs._state.update(queue=None, pid=None, dropped=0) # pylint: disable=protected-access + super().tearDown() + + def test_a_line_is_printed_by_the_writer_thread(self): + written = [] + with patch('core.capacity.logs.write', side_effect=written.append): + logs.emit({'event': 'ocl_capacity', 'decision': 'admitted', 'error': None}) + deadline = time.monotonic() + 2 + while not written and time.monotonic() < deadline: + time.sleep(0.01) + + self.assertEqual(written, ['{"event":"ocl_capacity","decision":"admitted"}']) + + def test_a_blocked_log_sink_never_blocks_the_call(self): + unblock = __import__('threading').Event() + written = [] + + def blocked_write(line): + unblock.wait(2) + written.append(line) + + with patch('core.capacity.logs.MAX_QUEUED_LINES', 2), patch('core.capacity.logs.write', blocked_write): + started = time.monotonic() + for number in range(10): + logs.emit({'n': number}) + self.assertLess(time.monotonic() - started, 0.5) + self.assertGreater(logs._state['dropped'], 0) # pylint: disable=protected-access + + unblock.set() + deadline = time.monotonic() + 2 + while logs._state['queue'].qsize() and time.monotonic() < deadline: # pylint: disable=protected-access + time.sleep(0.01) + logs.emit({'n': 'after'}) + while 'log_dropped' not in ''.join(written) and time.monotonic() < deadline: + time.sleep(0.01) + + self.assertEqual(json.loads(written[-1])['n'], 'after') + self.assertGreater(json.loads(written[-1])['log_dropped'], 0) + self.assertEqual(logs._state['dropped'], 0) # pylint: disable=protected-access + + def test_drain_prints_what_is_left(self): + out = StringIO() + with patch('core.capacity.logs.write_lines'): # no writer thread: the lines stay queued + logs.emit({'n': 1}) + with patch('sys.stdout', out): + logs.drain() + + self.assertEqual(out.getvalue(), '{"n":1}\n') diff --git a/core/settings.py b/core/settings.py index 273edab25..92388109a 100644 --- a/core/settings.py +++ b/core/settings.py @@ -685,6 +685,7 @@ def get_set_from_env(name): CAPACITY_CONFIG_CACHE_SECONDS = int(os.environ.get('CAPACITY_CONFIG_CACHE_SECONDS', 10)) CAPACITY_REDIS_TIMEOUT_SECONDS = float(os.environ.get('CAPACITY_REDIS_TIMEOUT_SECONDS', 0.5)) CAPACITY_REDIS_RETRY_SECONDS = int(os.environ.get('CAPACITY_REDIS_RETRY_SECONDS', 30)) +CAPACITY_REDIS_DEADLINE_SECONDS = float(os.environ.get('CAPACITY_REDIS_DEADLINE_SECONDS', 1.0)) CAPACITY_TASK_ID = os.environ.get('CAPACITY_TASK_ID', '') ANALYTICS_API = os.environ.get('ANALYTICS_API', 'http://host.docker.internal:8002') From 4be4acfe47c2204ca9de83c4fa80649df4e05060 Mon Sep 17 00:00:00 2001 From: Jonathan Payne Date: Tue, 29 Sep 2026 23:47:10 -0400 Subject: [PATCH 4/8] OpenConceptLab/ocl_online#275 | Capacity limit: fixes from Codex review, pass 3 - Redis operations run on a daemon thread with a short queue (core/capacity/threads.py) instead of a ThreadPoolExecutor, whose threads the interpreter waits for at exit: an operation stuck in DNS no longer holds up a gunicorn worker's recycle. When the queue is full, a call fails at once. - The log sink tracks lines in flight, and at exit waits for them for 2 s at most, so the last line (e.g. `manage.py capacity set`'s change) isn't lost, and a blocked sink can't hold the exit. The dropped-line count is kept under a lock. - The config refresh on the request path runs with a 500 ms statement timeout (CAPACITY_CONFIG_READ_TIMEOUT_MS), so a locked or stalled table can't hold a call; the last known config then applies, with enforce falling back to shadow. - Tests: process exit with a hung Redis call and with a blocked log sink (subprocesses), the last line printed at exit, concurrent emitters, a full Redis queue, and a locked config table. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TyAC9YAn5hP3yGSREkSN8m --- core/capacity/config.py | 16 +++- core/capacity/limiter.py | 39 ++++------ core/capacity/logs.py | 96 +++++++++++++++--------- core/capacity/tests/tests.py | 141 +++++++++++++++++++++++++++++------ core/capacity/threads.py | 41 ++++++++++ core/settings.py | 1 + 6 files changed, 252 insertions(+), 82 deletions(-) create mode 100644 core/capacity/threads.py diff --git a/core/capacity/config.py b/core/capacity/config.py index 574a1fbcf..22682bd44 100644 --- a/core/capacity/config.py +++ b/core/capacity/config.py @@ -142,6 +142,20 @@ def get_current(): return latest, resolve(latest.config if latest else None) +def read_current(): + """ + get_current() for the request path, bounded: a locked or stalled table costs the call at most + CAPACITY_CONFIG_READ_TIMEOUT_MS, then raises. Inside a transaction (tests) the timeout would outlive this read, + so it isn't set there; requests don't run in one. + """ + if connection.vendor != 'postgresql' or connection.in_atomic_block: + return get_current() + with transaction.atomic(): + with connection.cursor() as cursor: + cursor.execute('SET LOCAL statement_timeout = %s', [int(settings.CAPACITY_CONFIG_READ_TIMEOUT_MS)]) + return get_current() + + def get_config(): """The config in force, cached per process for CAPACITY_CONFIG_CACHE_SECONDS. Never raises.""" now = time.monotonic() @@ -152,7 +166,7 @@ def get_config(): if _cache['config'] is not None and time.monotonic() < _cache['expires_at']: return _cache['config'] # another thread just refreshed it try: - _, config = get_current() + _, config = read_current() except Exception as ex: # Keep the last good config (or the defaults) rather than fail the request, but never refuse on a # config that couldn't be confirmed: enforce falls back to shadow until a read succeeds. diff --git a/core/capacity/limiter.py b/core/capacity/limiter.py index 651d78ef7..f5c87a97a 100644 --- a/core/capacity/limiter.py +++ b/core/capacity/limiter.py @@ -20,7 +20,7 @@ import threading import time import uuid -from concurrent.futures import ThreadPoolExecutor, TimeoutError as FutureTimeoutError +from concurrent.futures import TimeoutError as FutureTimeoutError from contextlib import contextmanager import redis @@ -40,6 +40,7 @@ HEADER_DECISION, HEADER_LIMIT, HEADER_IN_FLIGHT, HEADER_TIER, HEADER_TIER_LIMIT, HEADER_TIER_IN_FLIGHT, HEADER_SUGGESTED_CONCURRENCY, CAPACITY_EXCEEDED_ERROR_CODE, LOG_EVENT, REDIS_KEY_PREFIX) from core.capacity.logs import emit +from core.capacity.threads import RedisThread from core.common.utils import get_event_metadata from core.users.constants import CORE_USER_GROUP, EARLY_ACCESS_GROUP @@ -138,8 +139,8 @@ class RedisLanes: scripts = {} unavailable_until = 0.0 lock = threading.Lock() - executor = None - executor_pid = None + redis_thread = None + redis_thread_pid = None @classmethod def use_client(cls, client): @@ -147,37 +148,29 @@ def use_client(cls, client): cls.client = client cls.scripts = {} cls.unavailable_until = 0.0 - if cls.executor is not None: - cls.executor.shutdown(wait=False, cancel_futures=True) - cls.executor = None + cls.redis_thread = None @classmethod - def get_executor(cls): - if cls.executor is None or cls.executor_pid != os.getpid(): + def get_redis_thread(cls): + if cls.redis_thread is None or cls.redis_thread_pid != os.getpid(): with cls.lock: - if cls.executor is None or cls.executor_pid != os.getpid(): - cls.executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix='ocl-capacity-redis') - cls.executor_pid = os.getpid() - return cls.executor + if cls.redis_thread is None or cls.redis_thread_pid != os.getpid(): # gunicorn forks workers + cls.redis_thread = RedisThread() + cls.redis_thread_pid = os.getpid() + return cls.redis_thread @classmethod def call(cls, operation, *args): """ - Run a Redis operation on this process's Redis thread, and wait at most CAPACITY_REDIS_DEADLINE_SECONDS for - it in all: socket timeouts don't bound DNS or a walk through the Sentinels. An operation still queued when - its caller gives up is skipped. One already running finishes, so an acquire that lands late holds a lease - nobody renews, which expires. + Run a Redis operation on this process's RedisThread, and wait at most CAPACITY_REDIS_DEADLINE_SECONDS for it + in all. An operation still queued when its caller gives up is skipped. One already running finishes, so an + acquire that lands late holds a lease nobody renews, which expires. """ - abandoned = threading.Event() - - def run(): - return None if abandoned.is_set() else operation(*args) - - future = cls.get_executor().submit(run) + future = cls.get_redis_thread().submit(operation, *args) try: return future.result(timeout=settings.CAPACITY_REDIS_DEADLINE_SECONDS) except FutureTimeoutError as ex: - abandoned.set() + future.cancel() raise redis.TimeoutError( f'no answer from Redis within {settings.CAPACITY_REDIS_DEADLINE_SECONDS} s') from ex diff --git a/core/capacity/logs.py b/core/capacity/logs.py index 54afcc949..be44a1242 100644 --- a/core/capacity/logs.py +++ b/core/capacity/logs.py @@ -6,45 +6,67 @@ import threading MAX_QUEUED_LINES = 1000 - -_state = {'queue': None, 'pid': None, 'dropped': 0} -_lock = threading.Lock() +DRAIN_SECONDS = 2 -def emit(record): +class LineSink: """ - Write one JSON object as one line of the API log, where CloudWatch metric filters match it. The API's other - timing lines are prints too: gunicorn captures stdout, and the `oclapi` logger has no handler in production. - - Never blocks the request: the line goes on a bounded queue that a writer thread prints, so a backed-up log sink - can't hold a call. When the queue is full the line is dropped, and the next line that gets through says how - many were (`log_dropped`). + A bounded queue of log lines and the daemon thread that prints them. Putting a line never blocks: when the queue + is full the line is dropped, and the next line that gets through says how many were (`log_dropped`). """ - dropped = _state['dropped'] - if dropped: - record = {**record, 'log_dropped': dropped} - line = json.dumps({key: value for key, value in record.items() if value is not None}, - separators=(',', ':'), default=str) - try: - get_queue().put_nowait(line) - _state['dropped'] -= dropped - except queue.Full: - _state['dropped'] += 1 + def __init__(self, max_lines): + self.lines = queue.Queue(maxsize=max_lines) + self.dropped = 0 + self.pending = 0 # queued or being written + self.idle = threading.Condition() + threading.Thread(target=self.run, name='ocl-capacity-log', daemon=True).start() + + def put(self, record): + with self.idle: + dropped = self.dropped + line = json.dumps({key: value for key, value in {**record, 'log_dropped': dropped or None}.items() + if value is not None}, separators=(',', ':'), default=str) + try: + self.lines.put_nowait(line) + except queue.Full: + self.dropped += 1 + return + self.dropped -= dropped + self.pending += 1 + + def run(self): + while True: + line = self.lines.get() + write(line) + with self.idle: + self.pending -= 1 + self.idle.notify_all() + + def drain(self, timeout): + """Wait until every line put so far is written, for `timeout` at most. True if they all were.""" + with self.idle: + return self.idle.wait_for(lambda: self.pending == 0, timeout=timeout) -def get_queue(): - if _state['pid'] != os.getpid(): +_sink = {'sink': None, 'pid': None} +_lock = threading.Lock() + + +def get_sink(): + if _sink['pid'] != os.getpid(): with _lock: - if _state['pid'] != os.getpid(): # first use in this process (gunicorn forks workers) - lines = queue.Queue(maxsize=MAX_QUEUED_LINES) - threading.Thread(target=write_lines, args=(lines,), name='ocl-capacity-log', daemon=True).start() - _state.update(queue=lines, pid=os.getpid(), dropped=0) - return _state['queue'] + if _sink['pid'] != os.getpid(): # first use in this process (gunicorn forks workers) + _sink.update(sink=LineSink(MAX_QUEUED_LINES), pid=os.getpid()) + return _sink['sink'] -def write_lines(lines): - while True: - write(lines.get()) +def emit(record): + """ + Write one JSON object as one line of the API log, where CloudWatch metric filters match it. The API's other + timing lines are prints too: gunicorn captures stdout, and the `oclapi` logger has no handler in production. + Never blocks the call: see LineSink. + """ + get_sink().put(record) def write(line): @@ -57,10 +79,10 @@ def write(line): @atexit.register def drain(): - """Print what's still queued when a worker exits (gunicorn recycles them).""" - lines = _state['queue'] if _state['pid'] == os.getpid() else None - while lines is not None: - try: - write(lines.get_nowait()) - except queue.Empty: - break + """ + When a process exits, give the writer DRAIN_SECONDS to finish the lines still queued or in flight: gunicorn + recycles workers, and `manage.py capacity` exits right after logging its change. + """ + sink = _sink['sink'] if _sink['pid'] == os.getpid() else None + if sink: + sink.drain(DRAIN_SECONDS) diff --git a/core/capacity/tests/tests.py b/core/capacity/tests/tests.py index ccc6c7377..085201b19 100644 --- a/core/capacity/tests/tests.py +++ b/core/capacity/tests/tests.py @@ -1,16 +1,21 @@ import json import os +import subprocess +import sys +import threading import time from io import StringIO from unittest.mock import patch, Mock import fakeredis +import psycopg2 import redis from django.conf import settings from django.contrib.auth.models import Group from django.core.exceptions import ValidationError from django.core.management import call_command, CommandError -from django.test import RequestFactory, override_settings +from django.db import connection +from django.test import RequestFactory, TransactionTestCase, override_settings from core.capacity.config import ( get_defaults, validate, merge, diff, resolve, save_config, get_config, clear_cache, _cache) @@ -24,6 +29,7 @@ RedisLanes, CapacityGate, LeaseRenewer, get_tier, get_queue_ms, lane_key, get_task_id, build_redis_client, get_status) from core.capacity.models import CapacityConfig +from core.capacity.threads import RedisThread from core.common.tests import OCLTestCase, OCLAPITestCase, PREVIEW_GROUP_NAME from core.common.utils import get_event_metadata from core.users.constants import CORE_USER_GROUP, EARLY_ACCESS_GROUP @@ -618,7 +624,7 @@ def hang(*_): self.assertFalse(RedisLanes.is_available()) def test_a_queued_redis_call_is_skipped_once_its_caller_gave_up(self): - release_first = __import__('threading').Event() + release_first = threading.Event() calls = [] def first(): @@ -975,24 +981,22 @@ def test_errors(self): class CapacityLogsTest(OCLTestCase): def setUp(self): super().setUp() - logs._state.update(queue=None, pid=None, dropped=0) # pylint: disable=protected-access + logs._sink.update(sink=None, pid=None) # pylint: disable=protected-access def tearDown(self): - logs._state.update(queue=None, pid=None, dropped=0) # pylint: disable=protected-access + logs._sink.update(sink=None, pid=None) # pylint: disable=protected-access super().tearDown() def test_a_line_is_printed_by_the_writer_thread(self): written = [] with patch('core.capacity.logs.write', side_effect=written.append): logs.emit({'event': 'ocl_capacity', 'decision': 'admitted', 'error': None}) - deadline = time.monotonic() + 2 - while not written and time.monotonic() < deadline: - time.sleep(0.01) + self.assertTrue(logs.get_sink().drain(2)) self.assertEqual(written, ['{"event":"ocl_capacity","decision":"admitted"}']) def test_a_blocked_log_sink_never_blocks_the_call(self): - unblock = __import__('threading').Event() + unblock = threading.Event() written = [] def blocked_write(line): @@ -1004,25 +1008,120 @@ def blocked_write(line): for number in range(10): logs.emit({'n': number}) self.assertLess(time.monotonic() - started, 0.5) - self.assertGreater(logs._state['dropped'], 0) # pylint: disable=protected-access + sink = logs.get_sink() + self.assertGreater(sink.dropped, 0) + self.assertFalse(sink.drain(0.1)) # a drain gives up at its deadline unblock.set() - deadline = time.monotonic() + 2 - while logs._state['queue'].qsize() and time.monotonic() < deadline: # pylint: disable=protected-access - time.sleep(0.01) + self.assertTrue(sink.drain(2)) logs.emit({'n': 'after'}) - while 'log_dropped' not in ''.join(written) and time.monotonic() < deadline: - time.sleep(0.01) + self.assertTrue(sink.drain(2)) self.assertEqual(json.loads(written[-1])['n'], 'after') self.assertGreater(json.loads(written[-1])['log_dropped'], 0) - self.assertEqual(logs._state['dropped'], 0) # pylint: disable=protected-access + self.assertEqual(sink.dropped, 0) - def test_drain_prints_what_is_left(self): - out = StringIO() - with patch('core.capacity.logs.write_lines'): # no writer thread: the lines stay queued + def test_drain_waits_for_a_line_being_written(self): + written = [] + + def slow_write(line): + time.sleep(0.2) + written.append(line) + + with patch('core.capacity.logs.write', slow_write): logs.emit({'n': 1}) - with patch('sys.stdout', out): - logs.drain() + time.sleep(0.05) # the writer has taken it off the queue and is writing it + logs.drain() + + self.assertEqual(written, ['{"n":1}']) + + def test_concurrent_emitters_keep_the_dropped_count_right(self): + unblock = threading.Event() + written = [] + + def blocked_write(line): + unblock.wait(2) + written.append(line) - self.assertEqual(out.getvalue(), '{"n":1}\n') + with patch('core.capacity.logs.MAX_QUEUED_LINES', 1), patch('core.capacity.logs.write', blocked_write): + threads = [threading.Thread(target=lambda: [logs.emit({'n': 1}) for _ in range(50)]) for _ in range(4)] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + unblock.set() + sink = logs.get_sink() + self.assertTrue(sink.drain(2)) + + # every line is either written, or counted as dropped: on a line that got through, or still pending + reported = sum(json.loads(line).get('log_dropped', 0) for line in written) + self.assertEqual(len(written) + reported + sink.dropped, 200) + self.assertGreaterEqual(sink.dropped, 0) + + +class CapacityRedisThreadTest(OCLTestCase): + def test_a_full_queue_fails_at_once(self): + unblock = threading.Event() + with patch.object(RedisThread, 'MAX_QUEUED', 1): + redis_thread = RedisThread() + redis_thread.submit(unblock.wait, 2) # running + time.sleep(0.05) + redis_thread.submit(lambda: None) # queued + + with self.assertRaises(redis.ConnectionError): + redis_thread.submit(lambda: None) + unblock.set() + + def test_a_process_exits_while_a_redis_call_hangs(self): + # A ThreadPoolExecutor's thread would hold the exit for the whole minute. + result = subprocess.run( + [sys.executable, '-c', + 'import time; from core.capacity.threads import RedisThread; RedisThread().submit(time.sleep, 60)'], + cwd=settings.BASE_DIR, capture_output=True, timeout=30, check=False) + + self.assertEqual(result.returncode, 0, result.stderr) + + def test_a_process_prints_its_last_line_before_it_exits(self): + result = subprocess.run( + [sys.executable, '-c', 'from core.capacity import logs; logs.emit({"n": 1})'], + cwd=settings.BASE_DIR, capture_output=True, text=True, timeout=30, check=False) + + self.assertEqual(result.stdout, '{"n":1}\n', result.stderr) + + def test_a_blocked_log_sink_doesnt_hold_up_exit(self): + code = ( + 'import sys, threading\n' + 'class Blocked:\n' + ' def write(self, _): threading.Event().wait()\n' + ' def flush(self): pass\n' + 'from core.capacity import logs\n' + 'sys.stdout = Blocked()\n' + 'logs.emit({"n": 1})\n' + ) + started = time.monotonic() + result = subprocess.run( + [sys.executable, '-c', code], cwd=settings.BASE_DIR, capture_output=True, timeout=30, check=False) + + self.assertEqual(result.returncode, 0, result.stderr) + self.assertLess(time.monotonic() - started, 20) + + +class CapacityConfigReadTimeoutTest(TransactionTestCase): + @override_settings(CAPACITY_CONFIG_READ_TIMEOUT_MS=200) + def test_a_locked_config_table_costs_a_call_at_most_the_read_timeout(self): + clear_cache() + locker = psycopg2.connect(**connection.get_connection_params()) + try: + locker.cursor().execute('LOCK TABLE capacity_configs IN ACCESS EXCLUSIVE MODE') + started = time.monotonic() + with self.assertLogs('oclapi', level='WARNING'): + config = get_config() + elapsed = time.monotonic() - started + finally: + locker.rollback() + locker.close() + clear_cache() + + self.assertLess(elapsed, 2) + self.assertEqual(config, get_defaults()) + self.assertEqual(get_config(), get_defaults()) # and the connection still works diff --git a/core/capacity/threads.py b/core/capacity/threads.py new file mode 100644 index 000000000..8925fe56e --- /dev/null +++ b/core/capacity/threads.py @@ -0,0 +1,41 @@ +"""The capacity limit's Redis thread. It imports nothing from Django, so a subprocess test can check process exit.""" +import queue +import threading +from concurrent.futures import Future + +import redis + + +class RedisThread: + """ + One daemon thread per process that runs the capacity limit's Redis operations, so that a caller can stop waiting + at a deadline: socket timeouts don't bound DNS or a walk through the Sentinels. + - It's a daemon, so an operation stuck in DNS never holds up a worker's exit. A ThreadPoolExecutor's threads + would: the interpreter waits for them at exit. + - Its queue is short. When it's full, a call fails at once instead of waiting behind a stuck operation. + - An operation whose caller gave up before it started is skipped. + """ + MAX_QUEUED = 8 + + def __init__(self): + self.jobs = queue.Queue(maxsize=self.MAX_QUEUED) + self.thread = threading.Thread(target=self.run, name='ocl-capacity-redis', daemon=True) + self.thread.start() + + def run(self): + while True: + future, operation, args = self.jobs.get() + if not future.set_running_or_notify_cancel(): + continue # cancelled: its caller stopped waiting before it started + try: + future.set_result(operation(*args)) + except Exception as ex: + future.set_exception(ex) + + def submit(self, operation, *args): + future = Future() + try: + self.jobs.put_nowait((future, operation, args)) + except queue.Full as ex: + raise redis.ConnectionError('the capacity Redis thread is busy') from ex + return future diff --git a/core/settings.py b/core/settings.py index 92388109a..daf1a486a 100644 --- a/core/settings.py +++ b/core/settings.py @@ -683,6 +683,7 @@ def get_set_from_env(name): # are runtime config (core.capacity.config); CAPACITY_LIMIT_MODE is only the default until staff change it. CAPACITY_LIMIT_MODE = os.environ.get('CAPACITY_LIMIT_MODE', 'shadow') CAPACITY_CONFIG_CACHE_SECONDS = int(os.environ.get('CAPACITY_CONFIG_CACHE_SECONDS', 10)) +CAPACITY_CONFIG_READ_TIMEOUT_MS = int(os.environ.get('CAPACITY_CONFIG_READ_TIMEOUT_MS', 500)) CAPACITY_REDIS_TIMEOUT_SECONDS = float(os.environ.get('CAPACITY_REDIS_TIMEOUT_SECONDS', 0.5)) CAPACITY_REDIS_RETRY_SECONDS = int(os.environ.get('CAPACITY_REDIS_RETRY_SECONDS', 30)) CAPACITY_REDIS_DEADLINE_SECONDS = float(os.environ.get('CAPACITY_REDIS_DEADLINE_SECONDS', 1.0)) From 2701b02e54f7e75b44022437724ec17a9c28b9b9 Mon Sep 17 00:00:00 2001 From: Jonathan Payne Date: Wed, 30 Sep 2026 23:51:58 -0400 Subject: [PATCH 5/8] OpenConceptLab/ocl_online#275 | Capacity limit: an acquire Redis runs too late takes no lease Found in a local soak with Redis behind Sentinel: while Redis hung (docker pause), acquire scripts that had timed out sat in the socket buffers and ran when Redis came back, taking leases nobody held. For up to lease_seconds after recovery the counts read high (in enforce mode that would refuse calls). The acquire now carries a not-after time (the caller's deadline plus 2 s for clock differences) and does nothing if Redis runs it later. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TyAC9YAn5hP3yGSREkSN8m --- core/capacity/limiter.py | 28 ++++++++++++++++++++++------ core/capacity/tests/tests.py | 15 +++++++++++++++ 2 files changed, 37 insertions(+), 6 deletions(-) diff --git a/core/capacity/limiter.py b/core/capacity/limiter.py index f5c87a97a..1d55b3a2a 100644 --- a/core/capacity/limiter.py +++ b/core/capacity/limiter.py @@ -54,9 +54,12 @@ end """ -# KEYS: one per lane. ARGV: token, lease ms, force ('1': take the lease even when a lane is full), then one limit -# per key. Returns {1 if the lease was taken else 0, then each lane's count before this call}. +# KEYS: one per lane. ARGV: token, lease ms, force ('1': take the lease even when a lane is full), not-after (epoch +# ms), then one limit per key. Returns {1 if the lease was taken else 0, then each lane's count before this call}, or +# {-1} if Redis runs it after its caller stopped waiting (it sat in a buffer while Redis hung): a lease taken then +# would be held by nobody. ACQUIRE_SCRIPT = _NOW_MS + """ +if now_ms > tonumber(ARGV[4]) then return {-1} end local lease_ms = tonumber(ARGV[2]) local result = {0} local full = false @@ -64,7 +67,7 @@ redis.call('ZREMRANGEBYSCORE', key, '-inf', now_ms) local count = redis.call('ZCARD', key) result[i + 1] = count - if count >= tonumber(ARGV[3 + i]) then full = true end + if count >= tonumber(ARGV[4 + i]) then full = true end end if ARGV[3] == '1' or not full then result[1] = 1 @@ -200,9 +203,17 @@ def mark_unavailable(cls): cls.unavailable_until = time.monotonic() + settings.CAPACITY_REDIS_RETRY_SECONDS @classmethod - def acquire(cls, keys, limits, token, lease_ms, force): # pylint: disable=too-many-arguments - """(whether the lease was taken, each lane's count before this call)""" - result = cls.call(cls.run_script, ACQUIRE_SCRIPT, keys, [token, lease_ms, '1' if force else '0', *limits]) + def acquire(cls, keys, limits, token, lease_ms, force, not_after_ms=None): # pylint: disable=too-many-arguments + """ + (whether the lease was taken, each lane's count before this call). Raises redis.TimeoutError if Redis ran it + after not_after_ms (by default, the deadline plus 2 s for clock differences), when its caller had given up. + """ + if not_after_ms is None: + not_after_ms = get_not_after_ms() + result = cls.call( + cls.run_script, ACQUIRE_SCRIPT, keys, [token, lease_ms, '1' if force else '0', not_after_ms, *limits]) + if int(result[0]) == -1: + raise redis.TimeoutError('Redis ran the acquire after its caller stopped waiting; no lease was taken') return bool(result[0]), [int(count) for count in result[1:]] @classmethod @@ -231,6 +242,11 @@ def scan(): return cls.call(scan) +def get_not_after_ms(): + """When an acquire stops being worth running: its caller's deadline, plus 2 s for clock differences.""" + return int((time.time() + settings.CAPACITY_REDIS_DEADLINE_SECONDS + 2) * 1000) + + def get_tier(user): """The user's highest plan tier, as capabilities resolve it: staff > core > early_access > preview.""" if user.is_staff or user.is_superuser: diff --git a/core/capacity/tests/tests.py b/core/capacity/tests/tests.py index 085201b19..b718baa71 100644 --- a/core/capacity/tests/tests.py +++ b/core/capacity/tests/tests.py @@ -598,6 +598,21 @@ def test_a_renewer_that_fails_to_start_still_releases(self): for key in gate.keys: self.assertEqual(self.redis.zcard(key), 0) + def test_an_acquire_that_redis_runs_too_late_takes_no_lease(self): + # Redis hung and ran the buffered script once it came back, after the caller had stopped waiting. + keys = [lane_key(LANE_API_HEAVY), lane_key(LANE_USER, 1)] + with self.assertRaises(redis.TimeoutError): + RedisLanes.acquire(keys, [4, 1], 'late', 60000, force=True, not_after_ms=int(time.time() * 1000) - 5000) + for key in keys: + self.assertEqual(self.redis.zcard(key), 0) + + with patch('core.capacity.limiter.get_not_after_ms', return_value=int(time.time() * 1000) - 5000): + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + self.assertEqual(gate.decision, DECISION_UNAVAILABLE) + self.assertFalse(gate.holding) + self.assertIn('no lease was taken', gate.error) + self.assertEqual(self.redis.zcard(lane_key(LANE_API_HEAVY)), 0) + def test_a_short_lease_never_shortens_a_lanes_ttl(self): self.redis.zadd(lane_key(LANE_API_HEAVY), {'long-lease': time.time() * 1000 + 600000}) self.redis.pexpire(lane_key(LANE_API_HEAVY), 1200000) From 2f399a1292bc48b0d2d8cdd06bc894c2ab7109b0 Mon Sep 17 00:00:00 2001 From: Jonathan Payne Date: Thu, 1 Oct 2026 00:04:59 -0400 Subject: [PATCH 6/8] OpenConceptLab/ocl_online#275 | Capacity limit: a refusal's scope is the paused lane when one is (Codex review of #922 + oclmap#86 together) Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TyAC9YAn5hP3yGSREkSN8m --- core/capacity/limiter.py | 10 ++++++++-- core/capacity/tests/tests.py | 13 +++++++++++++ 2 files changed, 21 insertions(+), 2 deletions(-) diff --git a/core/capacity/limiter.py b/core/capacity/limiter.py index 1d55b3a2a..5c3426255 100644 --- a/core/capacity/limiter.py +++ b/core/capacity/limiter.py @@ -339,6 +339,12 @@ def keys(self): def limits(self): return {lane: limit for lane, _, limit in self.lanes} + @property + def scope(self): + """The lane that refused the call (or would have): a paused one first, as it's what sets the long wait.""" + limits = self.limits + return next((lane for lane in self.full if limits[lane] == 0), self.full[0] if self.full else None) + @property def lease_ms(self): return self.config['lease_seconds'] * 1000 @@ -458,7 +464,7 @@ def record_error(self, ex): self.error = f'{ex.__class__.__name__}: {ex}'[:300] def get_refusal_response(self): - scope = self.full[0] + scope = self.scope return Response( { 'detail': f'Matching is busy. Please retry in {self.retry_after} seconds.', @@ -510,7 +516,7 @@ def lane_fields(lane, prefix): 'single_row': self.single_row, 'tier': self.tier, 'user_id': self.request.user.id, - 'scope': self.full[0] if self.full else None, + 'scope': self.scope, 'full': self.full or None, **lane_fields(LANE_API_HEAVY, ''), **lane_fields(LANE_API_HEAVY_TASK, 'task_'), diff --git a/core/capacity/tests/tests.py b/core/capacity/tests/tests.py index b718baa71..69c8de6c4 100644 --- a/core/capacity/tests/tests.py +++ b/core/capacity/tests/tests.py @@ -461,6 +461,19 @@ def test_a_paused_tier(self): self.assertEqual(gate.get_suggested_concurrency(), 0) self.assertTrue(gate.get_refusal_response().data['paused']) + def test_a_paused_lane_is_the_scope_reported(self): + self.configure(mode='enforce', enforce_for='all', reserve_single_row=0, + api_heavy={'cluster': 1, 'per_task': 5}, tiers={'preview': 0}) + self.acquire(make_user(CORE_USER_GROUP)) + + gate = self.acquire(make_user(PREVIEW_GROUP_NAME)) + + self.assertEqual(gate.full, [LANE_API_HEAVY, LANE_TIER]) + self.assertEqual(gate.scope, LANE_TIER) + self.assertEqual(gate.retry_after, 120) + data = gate.get_refusal_response().data + self.assertEqual((data['scope'], data['paused'], data['retry_after']), (LANE_TIER, True, 120)) + def test_retry_after_grows_with_the_calls_ahead_up_to_the_max(self): self.configure(tiers={'core': 10}, per_user={'core': 10}, api_heavy={'cluster': 10, 'per_task': 10}) user = make_user(CORE_USER_GROUP) From 040a689188cf7258f800461cea53723a2f1ff6cd Mon Sep 17 00:00:00 2001 From: Jonathan Payne Date: Thu, 1 Oct 2026 00:53:52 -0400 Subject: [PATCH 7/8] OpenConceptLab/ocl_online#275 | Capacity limit: skip Redis for 15 s after an error, not 30, so one error can't cost a long call its lease (independent review of #922 + oclmap#86) Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TyAC9YAn5hP3yGSREkSN8m --- core/capacity/tests/tests.py | 4 ++++ core/settings.py | 4 +++- 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/core/capacity/tests/tests.py b/core/capacity/tests/tests.py index 69c8de6c4..216af6bf8 100644 --- a/core/capacity/tests/tests.py +++ b/core/capacity/tests/tests.py @@ -520,6 +520,10 @@ def test_renewal_extends_every_lease_or_none(self): gate.renew() self.assertFalse(self.redis.exists(key)) # a lost lease isn't re-added + def test_one_redis_error_cant_cost_a_call_two_renewals(self): + from core.capacity.config import DEFAULTS + self.assertLess(settings.CAPACITY_REDIS_RETRY_SECONDS, DEFAULTS['lease_seconds'] - 2 * DEFAULTS['renew_seconds']) + def test_the_renewer_stops_once_a_lease_is_lost(self): gate = Mock(config={'renew_seconds': 0.01, 'max_hold_seconds': 60}, lease_lost=True) renewer = LeaseRenewer(gate) diff --git a/core/settings.py b/core/settings.py index daf1a486a..8daeeefa6 100644 --- a/core/settings.py +++ b/core/settings.py @@ -685,7 +685,9 @@ def get_set_from_env(name): CAPACITY_CONFIG_CACHE_SECONDS = int(os.environ.get('CAPACITY_CONFIG_CACHE_SECONDS', 10)) CAPACITY_CONFIG_READ_TIMEOUT_MS = int(os.environ.get('CAPACITY_CONFIG_READ_TIMEOUT_MS', 500)) CAPACITY_REDIS_TIMEOUT_SECONDS = float(os.environ.get('CAPACITY_REDIS_TIMEOUT_SECONDS', 0.5)) -CAPACITY_REDIS_RETRY_SECONDS = int(os.environ.get('CAPACITY_REDIS_RETRY_SECONDS', 30)) +# How long a process skips Redis after an error. Under lease_seconds - 2 x renew_seconds (60 - 40), so one error +# can't make a long call miss two renewals and lose its lease. +CAPACITY_REDIS_RETRY_SECONDS = int(os.environ.get('CAPACITY_REDIS_RETRY_SECONDS', 15)) CAPACITY_REDIS_DEADLINE_SECONDS = float(os.environ.get('CAPACITY_REDIS_DEADLINE_SECONDS', 1.0)) CAPACITY_TASK_ID = os.environ.get('CAPACITY_TASK_ID', '') From c5282f4a341a375575cd9c9ba42ff45d661a442f Mon Sep 17 00:00:00 2001 From: Jonathan Payne Date: Thu, 1 Oct 2026 00:55:42 -0400 Subject: [PATCH 8/8] OpenConceptLab/ocl_online#275 | Capacity limit tests: wrap a long line Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01TyAC9YAn5hP3yGSREkSN8m --- core/capacity/tests/tests.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/core/capacity/tests/tests.py b/core/capacity/tests/tests.py index 216af6bf8..0e1699b15 100644 --- a/core/capacity/tests/tests.py +++ b/core/capacity/tests/tests.py @@ -522,7 +522,8 @@ def test_renewal_extends_every_lease_or_none(self): def test_one_redis_error_cant_cost_a_call_two_renewals(self): from core.capacity.config import DEFAULTS - self.assertLess(settings.CAPACITY_REDIS_RETRY_SECONDS, DEFAULTS['lease_seconds'] - 2 * DEFAULTS['renew_seconds']) + self.assertLess( + settings.CAPACITY_REDIS_RETRY_SECONDS, DEFAULTS['lease_seconds'] - 2 * DEFAULTS['renew_seconds']) def test_the_renewer_stops_once_a_lease_is_lost(self): gate = Mock(config={'renew_seconds': 0.01, 'max_hold_seconds': 60}, lease_lost=True)