Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 9 additions & 5 deletions src/schematic/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -299,7 +299,7 @@ def _build_preflight(options: Optional[CheckFlagOptions]) -> Optional[PreflightR
def _preflight_quantity(usage: float) -> int:
"""Cast a usage onto the integer the preflight body carries.

A hold can be sized from a fractional usage, but the API's preflight usage
A caller can pass a fractional usage, but the API's preflight usage
is an integer. A preflight asks an upper-bound question ("would this action
be allowed?"), so a fraction rounds up: the check must not pass on less
usage than the operation is about to record.
Expand Down Expand Up @@ -935,7 +935,7 @@ def failure(reason: str) -> CheckResult:
flag_key, options, reason, self._resolve_default(flag_key, _check_options_to_flag_options(options)),
)

if not _is_valid_quantity(options.usage):
if options.usage is None or not _is_valid_quantity(options.usage):
self.logger.error(
f"Server reservation: invalid usage {options.usage!r} for flag {flag_key}; "
"must be a finite, non-negative number"
Expand All @@ -953,7 +953,9 @@ def failure(reason: str) -> CheckResult:
flag_key,
company=company,
user=user,
quantity=options.usage,
# Whole units, like the local lease path: the settle bills
# ceil(actual), so a fractional hold would come up short.
quantity=_preflight_quantity(options.usage),
expires_at=dt.datetime.now(dt.timezone.utc) + dt.timedelta(seconds=self._reservation_ttl),
**_reservation_request_kwargs(options),
)
Expand Down Expand Up @@ -1778,7 +1780,7 @@ def failure(reason: str) -> CheckResult:
flag_key, options, reason, self._resolve_default(flag_key, _check_options_to_flag_options(options)),
)

if not _is_valid_quantity(options.usage):
if options.usage is None or not _is_valid_quantity(options.usage):
self.logger.error(
f"Server reservation: invalid usage {options.usage!r} for flag {flag_key}; "
"must be a finite, non-negative number"
Expand All @@ -1796,7 +1798,9 @@ def failure(reason: str) -> CheckResult:
flag_key,
company=company,
user=user,
quantity=options.usage,
# Whole units, like the local lease path: the settle bills
# ceil(actual), so a fractional hold would come up short.
quantity=_preflight_quantity(options.usage),
expires_at=dt.datetime.now(dt.timezone.utc) + dt.timedelta(seconds=self._reservation_ttl),
**_reservation_request_kwargs(options),
)
Expand Down
17 changes: 11 additions & 6 deletions src/schematic/leases/check.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,10 @@ async def check_with_lease(
log = deps.logger
mode = options.on_acquire_failure or "fail-closed"
usage = options.usage
# One budget for the whole check. Each wait on another caller's acquire or
# extend is capped against this, not against a fresh timeout, so a check
# cannot take its timeout once per step.
deadline = None if options.timeout is None else time.monotonic() + options.timeout

# A malformed usage must never reach the stores, and the caller asked for a
# contract for exactly this case, so resolve it through that rather than
Expand Down Expand Up @@ -206,7 +210,7 @@ async def failure(reason: str) -> "CheckResult":

# The caller's per-check timeout governs the lease wire calls, the same way
# it governs the plain check's.
lease = await deps.manager.acquire_if_needed(resolved_company.id, credit_id, options.timeout)
lease = await deps.manager.acquire_if_needed(resolved_company.id, credit_id, options.timeout, deadline)
if lease is None:
return await failure("lease_acquire_failed")

Expand All @@ -221,7 +225,7 @@ async def failure(reason: str) -> "CheckResult":
if reserve is None:
# Pass the cost as required_credits so a single large request
# extends even while the ratio sits above the water mark.
await deps.manager.maybe_extend(resolved_company.id, credit_id, credit_cost, options.timeout)
await deps.manager.maybe_extend(resolved_company.id, credit_id, credit_cost, options.timeout, deadline)
reserve = await deps.lease_store.try_reserve(resolved_company.id, credit_id, credit_cost)
except Exception as err:
log.error(f"Lease check: reserve against {resolved_company.id}/{credit_id} failed: {err}")
Expand Down Expand Up @@ -259,11 +263,12 @@ async def failure(reason: str) -> "CheckResult":
# claims whatever slice of the add landed and refunds it; a None says
# nothing landed, so refund the debit directly. Both are pinned to the
# lease the debit landed on (the record carries that id, so consume
# pins to it too), never to the acquired one. If the undo itself fails,
# accept the bounded leak: the slice comes back at lease expiry, which
# beats risking a double refund.
# pins to it too), never to the acquired one; a debit that names no
# lease is not refunded at all, as consume would not either. If the
# undo itself fails, accept the bounded leak: the slice comes back at
# lease expiry, which beats risking a double refund.
try:
if await deps.reservations.consume(record.id, 0) is None:
if await deps.reservations.consume(record.id, 0) is None and reserve.lease_id:
await deps.lease_store.refund(resolved_company.id, credit_id, credit_cost, reserve.lease_id)
except Exception as undo_err:
log.warning(
Expand Down
58 changes: 46 additions & 12 deletions src/schematic/leases/lease_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -189,13 +189,18 @@ def sweep_interval(self) -> float:
return self._config.sweep_interval or DEFAULT_SWEEP_INTERVAL

async def acquire_if_needed(
self, company_id: str, credit_type_id: str, timeout: Optional[float] = None
self,
company_id: str,
credit_type_id: str,
timeout: Optional[float] = None,
deadline: Optional[float] = None,
) -> Optional[LeaseState]:
"""The slot's live lease, acquiring one over the wire if none is live.

``timeout`` governs the wire call this caller starts. A caller that
joins an in-flight acquire rides the first caller's timeout, since
there is one shared call to time out.
joins an in-flight acquire waits on the first caller's call, but only
until ``deadline`` (a ``time.monotonic()`` instant), and then resolves
to no lease while the call runs on for everybody else.
"""
try:
existing = await self._lease_store.get(company_id, credit_type_id)
Expand All @@ -214,7 +219,15 @@ async def acquire_if_needed(
key = lease_key(company_id, credit_type_id)
inflight = self._inflight_acquire.get(key)
if inflight is not None:
return await asyncio.shield(inflight.task)
joined = await self._join_within(inflight.task, deadline)
if joined is _JOIN_TIMED_OUT:
logger.debug(
"Acquire in flight for %s/%s outlasted the caller's deadline; not waiting on it",
company_id,
credit_type_id,
)
return None
return joined
return await self._single_flight(
self._inflight_acquire, key, self._acquire(company_id, credit_type_id, timeout)
)
Expand Down Expand Up @@ -269,6 +282,7 @@ async def maybe_extend(
credit_type_id: str,
required_credits: Optional[float] = None,
timeout: Optional[float] = None,
deadline: Optional[float] = None,
) -> Optional[LeaseState]:
"""Extend the slot's lease when the local view warrants it.

Expand All @@ -283,21 +297,29 @@ async def maybe_extend(
and fail its post-extend retry with credits still sitting on the
server. A flight it finds on the way back is only joined if that one
covers the shortfall too; a smaller one is waited out, never inherited.

Waits on other callers' flights end at ``deadline`` (a
``time.monotonic()`` instant), or ``timeout`` from now without one.
"""
return await self._maybe_extend(company_id, credit_type_id, required_credits, timeout)
return await self._maybe_extend(company_id, credit_type_id, required_credits, timeout, deadline)

async def _maybe_extend(
self,
company_id: str,
credit_type_id: str,
required_credits: Optional[float],
timeout: Optional[float],
deadline: Optional[float] = None,
) -> Optional[LeaseState]:
# A joiner waits on someone else's wire call, which runs on whatever
# timeout ITS caller set (a background refresh uses the client
# default). So the wait is capped at this caller's own timeout: a check
# with 200ms to spend must not sit behind a 30s extend.
join_deadline = None if timeout is None else time.monotonic() + timeout
# with 200ms to spend must not sit behind a 30s extend. A check passes
# the deadline it set when it started, so the time it already spent
# acquiring and reserving comes out of the same 200ms.
join_deadline = deadline
if join_deadline is None and timeout is not None:
join_deadline = time.monotonic() + timeout
# Joins are budgeted, extends of our own are not: a caller may wait out
# flights that ask for too little, but once the budget runs out it
# issues its own single extend rather than joining again. Without the
Expand Down Expand Up @@ -341,12 +363,24 @@ async def _maybe_extend(
# The flight asked for at least what we need: every
# watermark-driven joiner, and any check the tranche covers.
# One wire call serves all of them, which is the point of
# single-flight.
# single-flight. But the flight re-checks against its starter's
# need, not ours: if a sibling's extend landed first it may
# have skipped the wire call and left less than we need, so
# hold its result to our own requirement before taking it.
# Only the requirement: a server that granted less than asked
# can leave the slot under the water mark, and re-extending
# for that would turn every water-mark joiner into a wire call.
if additional_amount <= (inflight.requested_additional or 0.0):
return joined
# It asked for less. Go round again to re-read the slot it just
# moved, so what we ask for next is sized against the balance
# it left rather than the one we started from.
if (
joined is None
or required_credits is None
or joined.local_remaining_credits >= required_credits
):
return joined
# It asked for less, or left us short. Go round again to
# re-read the slot it just moved, so what we ask for next is
# sized against the balance it left rather than the one we
# started from.
joins_left -= 1
continue
return await self._single_flight(
Expand Down
8 changes: 6 additions & 2 deletions src/schematic/leases/redis_reservation_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,12 +181,16 @@ async def consume(self, reservation_id: str, credits_consumed: float) -> Optiona

consumed = clamp_consumption(credits_consumed, reserved)
refund = reserved - consumed
if refund > 0:
lease_id = raw.get("leaseId")
# A hold with no lease id cannot be pinned, and an empty pin disables
# the lease check, so the refund would land on whatever lease holds the
# slot now. Skip it: the slice comes back when its lease expires.
if refund > 0 and lease_id:
# The lease store owns the lease hash, which keeps this cross-key
# write out of a single Lua script. Pinned to the reservation's
# lease so a hold carved out of an expired lease cannot inflate a
# successor's balance.
await self._lease_store.refund(company_id, credit_type_id, refund, raw.get("leaseId"))
await self._lease_store.refund(company_id, credit_type_id, refund, lease_id)
return consumed

async def reserved_credits(self, company_id: str, credit_type_id: str) -> float:
Expand Down
7 changes: 6 additions & 1 deletion src/schematic/leases/reservation_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,11 @@ async def consume(self, reservation_id: str, credits_consumed: float) -> Optiona
remainder is refunded to the lease (pinned to the reservation's lease),
and the clamped figure is returned. A crash between the claim and the
refund loses the refund, never double-refunds.

A reservation with no lease id is claimed but not refunded: with
nothing to pin to, the refund would land on whichever lease holds the
slot now and could inflate a successor. The slice comes back when its
lease expires. Every SDK on a shared Redis must agree on this.
"""

@abc.abstractmethod
Expand Down Expand Up @@ -84,7 +89,7 @@ async def consume(self, reservation_id: str, credits_consumed: float) -> Optiona
return None
consumed = clamp_consumption(credits_consumed, reservation.credits_reserved)
refund = reservation.credits_reserved - consumed
if refund > 0:
if refund > 0 and reservation.lease_id:
await self._lease_store.refund(
reservation.company_id,
reservation.credit_type_id,
Expand Down
21 changes: 11 additions & 10 deletions tests/custom/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -1800,13 +1800,14 @@ def test_mints_a_fresh_idempotency_key_per_check(self):
second = self.schematic.features.check_and_reserve_flag.call_args.kwargs["idempotency_key"]
self.assertNotEqual(first, second)

def test_a_fractional_usage_sizes_the_hold_and_rounds_the_preflight_up(self):
self.schematic.check("inference", company={"id": "co_1"}, options=CheckOptions(usage=0.5))
def test_a_fractional_usage_rounds_the_hold_and_the_preflight_up(self):
self.schematic.check("inference", company={"id": "co_1"}, options=CheckOptions(usage=2.5))
kwargs = self.schematic.features.check_and_reserve_flag.call_args.kwargs
self.assertEqual(kwargs["quantity"], 0.5)
# The hold takes the fraction; the preflight's usage is an integer, and
# rounding it down would ask about less usage than is about to land.
self.assertEqual(kwargs["preflight"], PreflightRequestBody(usage=1))
# The settle bills whole units, so a 2.5 hold would take 2.5 * rate and
# the track would bill 3 * rate. Rounding the preflight down would ask
# about less usage than is about to land.
self.assertEqual(kwargs["quantity"], 3)
self.assertEqual(kwargs["preflight"], PreflightRequestBody(usage=3))

def test_an_integral_float_usage_reaches_the_preflight_unchanged(self):
self.schematic.check(
Expand Down Expand Up @@ -2268,11 +2269,11 @@ async def test_mints_a_fresh_idempotency_key_per_check(self):
second = self.client.features.check_and_reserve_flag.call_args.kwargs["idempotency_key"]
assert first != second

async def test_a_fractional_usage_sizes_the_hold_and_rounds_the_preflight_up(self):
await self.client.check("inference", company={"id": "co_1"}, options=CheckOptions(usage=0.5))
async def test_a_fractional_usage_rounds_the_hold_and_the_preflight_up(self):
await self.client.check("inference", company={"id": "co_1"}, options=CheckOptions(usage=2.5))
kwargs = self.client.features.check_and_reserve_flag.call_args.kwargs
assert kwargs["quantity"] == 0.5
assert kwargs["preflight"] == PreflightRequestBody(usage=1)
assert kwargs["quantity"] == 3
assert kwargs["preflight"] == PreflightRequestBody(usage=3)

async def test_a_reservation_ttl_above_the_cap_is_clamped(self):
client = _async_server_client(credit_leases=CreditLeaseConfig(default_reservation_ttl=7200.0))
Expand Down
Loading
Loading