Files
treg/tests/test_asynctasks.py
SToneX 2e3097cb1c feat(call): spool large metered answers and settle from their usage
A metered answer is buffered whole before headers so settlement sees complete evidence, and
anything over 8 MiB fails uncharged. Providers that inline generated media in JSON (Gemini
returns a base64 image plus a large thoughtSignature: ~9 MB at 2K, ~23 MB at 4K) can never fit.
Streaming instead would settle after the response starts, losing the exact X-Treg-Cost-Micro
and adding a second close-once path for disconnects, routed children and overflow; raising the
buffer would hold tens of megabytes per call in RAM.

An endpoint now declares `spooled_response: true`. Its metered 2xx is written to an unlinked
temp file under `spool_max_bytes` (64 MiB) and a per-process `spool_budget_bytes` (512 MiB),
claimed whole up front when the provider declares a Content-Length;
crossing either is the existing uncharged `response_buffer_limit`. The file is parsed once with
stdlib json in a worker thread (a 23 MB answer: ~30 ms, ~45 MB peak) behind
`spool_parse_concurrency`, and only the top-level keys its usage paths start at
(`settlement.usage_roots`) become the body every settlement consumer reads, so the evidence
cannot drift from the price. Settlement runs before headers as before; the router then relays
the file byte for byte, and its close (or garbage collection) returns the budget once. Spooled
bodies are not archived or kept for idempotent replay. Errors, own-key calls and routed
children keep their existing paths. `_buffer_response` and the spool share the refusal text and
the content-length rewrite.

The validator requires settle: usage and refuses the field beside async, resource_ownership or
managed_resource. AGENTS.md non-negotiable 4 records the new path.

Fragments: architecture/proxy-model.md, money.md, archive.md, catalog.md, interface/api.md,
ops/deploy.md.
2026-09-29 23:27:00 +08:00

1507 lines
72 KiB
Python

"""Deferred settlement for asynchronous metered catalog calls."""
from __future__ import annotations
from conftest import verified_signup
import asyncio
import json
import re
from datetime import timedelta
import pytest
from httpx import AsyncClient, ReadTimeout
from sqlalchemy import event
from sqlmodel import select
from treg.application import asynctasks as task_app
from treg.application.call import service as call_service
from treg.application.call.resolve import MarketplaceCall
from treg.application.call.types import UpstreamResponse
from treg.config import get_settings
from treg import archive, audit, reconcile
from treg.domain import asynctasks
from treg.domain.catalog import store as catalog_store
from treg.domain import money as ledger
from treg.domain.money import settlement
from treg.infra.db import _engine, session_maker
from treg.models import (
ArchiveKey, ArchiveSnapshot, AsyncResourceRecord, AsyncTaskRecord, CallRecord, Hold, LedgerEntry,
Tool,
)
from treg.timeutil import utcnow_naive
EP = "replicate.image-gen.flux-schnell"
def test_all_generation_catalog_entries_forbid_cache_including_extended():
entries = [ep for ep in catalog_store.load().endpoints
if ep["platform"] in {"image-gen", "video-gen", "voice-gen"}]
assert entries
for ep in entries:
assert ep["cache"] == "forbidden", ep["id"]
assert not archive.storable(ep), ep["id"]
@pytest.mark.parametrize("overrides", [
{"async_owner_call_id": None}, # Other free utilities still follow ordinary policy.
{"cost_type": "per_call"}, # A zero estimate does not prove a paid endpoint is free.
{"estimate_micro": 1},
{"tier": "platform-overflow"},
{"billed_oauth": True},
])
def test_poll_money_exception_requires_an_owned_explicitly_free_platform_read(overrides):
fields = dict(tool=Tool(org_id=1, name="poll", owner="test@example.invalid",
base_url="https://example.invalid", host="example.invalid"),
upstream="https://example.invalid/task", consumed=set(),
endpoint_id="replicate.predictions.get", provider="replicate", tier="platform",
cost_type="free", async_owner_call_id="original-submission")
mk = MarketplaceCall(**(fields | overrides))
assert mk.metered
assert not mk.free_owned_poll
def _response(status: int, document: object) -> UpstreamResponse:
body = json.dumps(document).encode()
async def stream():
yield body
async def close():
return None
return UpstreamResponse(status, ((b"content-type", b"application/json"),), stream(), close)
@pytest.fixture
def replicate_platform(monkeypatch):
monkeypatch.setenv("TREG_PLATFORM_KEY_REPLICATE", "test-platform-token")
monkeypatch.setenv("TREG_PLATFORM_PROVIDERS", "replicate")
get_settings.cache_clear()
yield
get_settings.cache_clear()
async def _submit(clients: AsyncClient, monkeypatch, document: dict):
async def fake_relay(*args, **kwargs):
return _response(201, document)
monkeypatch.setattr(call_service, "relay", fake_relay)
return await clients.post(f"/call/{EP}", json={"input": {
"prompt": "A red kite over a beach.", "num_outputs": 1,
"aspect_ratio": "1:1", "output_format": "webp",
}})
@pytest.mark.parametrize("legacy_cache", [False, True])
async def test_generation_is_never_replayed_across_orgs(
clients: AsyncClient, monkeypatch, replicate_platform, legacy_cache,
):
monkeypatch.setattr(get_settings(), "archive_mode", "serve")
monkeypatch.setattr(get_settings(), "archive_serve_endpoints", EP)
monkeypatch.setattr(get_settings(), "archive_serve_percent", 100)
entry = catalog_store.load().by_id[EP]
if legacy_cache:
monkeypatch.setitem(entry, "cache", "transient")
first = await _submit(clients, monkeypatch, {"id": "private-first-task"})
assert first.status_code == 201
await archive.drain()
monkeypatch.setitem(entry, "cache", "forbidden")
other = await verified_signup(clients, json={"email": "cache-stranger@example.com"})
async def live(*args, **kwargs):
return _response(201, {"id": "private-second-task"})
monkeypatch.setattr(call_service, "relay", live)
second = await clients.post(f"/call/{EP}", json={"input": {
"prompt": "A red kite over a beach.", "num_outputs": 1,
"aspect_ratio": "1:1", "output_format": "webp",
}}, headers={"X-Treg-Token": other.json()["token"]})
assert second.status_code == 201
assert second.json()["id"] == "private-second-task"
assert second.headers.get("X-Treg-Cache") != "hit"
async def test_worker_errors_grow_backoff_and_progress_resets_it(
clients: AsyncClient, monkeypatch, replicate_platform,
):
call_id = await _due_submission(clients, monkeypatch, {})
async def failing(row, client):
return 500, b'{}'
monkeypatch.setattr(task_app, "_poll", failing)
for failure in range(1, 5):
before = utcnow_naive()
await task_app.settle_due()
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
assert row.consecutive_failures == failure
assert (row.next_check_at - before).total_seconds() >= min(900, 120 * 2 ** (failure - 1))
row.next_check_at = utcnow_naive() - timedelta(seconds=1)
await db.commit()
async def progress(row, client):
return 200, b'{"status":"processing"}'
monkeypatch.setattr(task_app, "_poll", progress)
await task_app.settle_due()
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
assert row.consecutive_failures == 0
assert await db.get(Hold, call_id) is not None
async def test_queued_worker_rows_are_not_claimed_before_a_poll_slot(
clients: AsyncClient, monkeypatch, replicate_platform,
):
await _due_submission(clients, monkeypatch, {})
second_id = await _due_submission(clients, monkeypatch, {})
monkeypatch.setattr(task_app, "PROVIDER_CONCURRENCY", 1)
started, release = asyncio.Event(), asyncio.Event()
async def slow(row, client):
started.set()
await release.wait()
return 200, b'{"status":"processing"}'
monkeypatch.setattr(task_app, "_poll", slow)
tick = asyncio.create_task(task_app.settle_due())
try:
await asyncio.wait_for(started.wait(), 5)
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, second_id)
assert row.attempts == 0
finally:
release.set()
await tick
# POLL_TIMEOUT_S wraps only the upstream poll, so it can be near-zero. PROCESS_TIMEOUT_S also
# wraps the claim's own DB round trip: at 10 ms a loaded CI runner fires it BEFORE the claim,
# which is correctly reported as backed_off with the row untouched - and then this test, which
# wants the post-claim path, fails on `consecutive_failures == 1`. Give the claim room; the hung
# poll is what the deadline must cut, and it never returns regardless.
@pytest.mark.parametrize("deadline, seconds", [("POLL_TIMEOUT_S", 0.01), ("PROCESS_TIMEOUT_S", 0.5)])
async def test_worker_bounds_whole_poll_and_keeps_hold_on_timeout(
clients: AsyncClient, monkeypatch, replicate_platform, deadline, seconds,
):
call_id = await _due_submission(clients, monkeypatch, {})
monkeypatch.setattr(task_app, deadline, seconds)
async def hangs(row, client):
await asyncio.Event().wait()
monkeypatch.setattr(task_app, "_poll", hangs)
result = await asyncio.wait_for(task_app.settle_due(), 5)
assert result.backed_off == 1
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
assert row.status == "pending" and row.consecutive_failures == 1
assert await db.get(Hold, call_id) is not None
async def test_claim_wait_is_inside_processing_deadline(
clients: AsyncClient, monkeypatch, replicate_platform,
):
call_id = await _due_submission(clients, monkeypatch, {})
monkeypatch.setattr(task_app, "PROCESS_TIMEOUT_S", 0.01)
async def blocked_claim(*args):
await asyncio.Event().wait()
monkeypatch.setattr(task_app, "_claim_due", blocked_claim)
result = await asyncio.wait_for(task_app.settle_due(), 5)
assert result.claimed == 0 and result.backed_off == 1
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
assert row.attempts == 0 and row.consecutive_failures == 0
assert await db.get(Hold, call_id) is not None
async def test_expired_worker_cannot_replace_new_claim_backoff(
clients: AsyncClient, monkeypatch, replicate_platform,
):
call_id = await _due_submission(clients, monkeypatch, {})
started, release = asyncio.Event(), asyncio.Event()
polls = 0
async def poll(row, client):
nonlocal polls
polls += 1
if polls == 1:
async with session_maker() as db:
live = await db.get(AsyncTaskRecord, call_id)
live.next_check_at = utcnow_naive() - timedelta(seconds=1)
await db.commit()
started.set()
await release.wait()
return 200, b'{"status":"processing"}'
return 500, b'{}'
monkeypatch.setattr(task_app, "_poll", poll)
first = asyncio.create_task(task_app.settle_due())
try:
await asyncio.wait_for(started.wait(), 5)
await task_app.settle_due()
async with session_maker() as db:
due = (await db.get(AsyncTaskRecord, call_id)).next_check_at
finally:
release.set()
await first
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
assert row.attempts == 2 and row.consecutive_failures == 1
assert row.next_check_at == due
async def test_inline_success_waits_for_usage_and_then_closes_original_hold(
clients: AsyncClient, monkeypatch, replicate_platform,
):
submitted = await _submit(clients, monkeypatch, {"id": "usage-late"})
call_id = submitted.headers["X-Treg-Call-Id"]
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
row.settlement_basis = settlement.derive_basis(
{"settle": "usage", "usage": {"path": "usage.cost", "unit": "usd"},
"fallback": {"value": 0.003}, "type": "per_success", "currency": "USD"},
request={}, input_schema={}, unit_micro=1_000_000, terminal=True)
await db.commit()
document = {"status": "succeeded", "output": ["https://example.invalid/result.png"]}
async def live(*args, **kwargs):
return _response(200, document)
monkeypatch.setattr(call_service, "relay", live)
assert (await clients.get("/call/replicate.predictions.get?id=usage-late")).status_code == 200
async with session_maker() as db:
assert (await db.get(AsyncTaskRecord, call_id)).status == "pending"
assert await db.get(Hold, call_id) is not None
document["usage"] = {"cost": 0.002}
assert (await clients.get("/call/replicate.predictions.get?id=usage-late")).status_code == 200
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
assert row.status == "settled" and row.settled_micro == 2000
assert await db.get(Hold, call_id) is None
async def test_inline_settlement_failure_preserves_response_and_cron_recovers(
clients: AsyncClient, monkeypatch, replicate_platform,
):
call_id = await _due_submission(clients, monkeypatch, {"status": "succeeded"})
real_finish = task_app._finish
async def unavailable(*args, **kwargs):
raise RuntimeError("temporary database failure")
monkeypatch.setattr(task_app, "_finish", unavailable)
async def live(*args, **kwargs):
return _response(200, {"status": "succeeded"})
monkeypatch.setattr(call_service, "relay", live)
response = await clients.get("/call/replicate.predictions.get?id=prediction-worker")
assert response.status_code == 200 and response.json() == {"status": "succeeded"}
assert response.headers["X-Treg-Cost-Micro"] == "0"
async with session_maker() as db:
assert (await db.get(AsyncTaskRecord, call_id)).status == "pending"
assert await db.get(Hold, call_id) is not None
monkeypatch.setattr(task_app, "_finish", real_finish)
assert (await task_app.settle_due()).settled == 1
async def test_settle_fork_keeps_hold_and_writes_pending_row(
clients: AsyncClient, monkeypatch, replicate_platform,
):
response = await _submit(clients, monkeypatch, {
"id": "prediction-1", "urls": {"get": "https://api.replicate.com/v1/predictions/1"}})
assert response.status_code == 201
assert response.headers["X-Treg-Cost-Micro"] == "3000"
call_id = response.headers["X-Treg-Call-Id"]
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
hold = await db.get(Hold, call_id)
assert row is not None and hold is not None
assert (row.task_id, row.poll_url, row.status) == ("prediction-1", None, "pending")
assert row.settlement_basis["when"] == "terminal"
@pytest.mark.parametrize("status", [200, 500])
async def test_owned_free_poll_creates_no_money_entries(
clients: AsyncClient, monkeypatch, replicate_platform, status: int,
):
submitted = await _submit(clients, monkeypatch, {"id": "free-poll-task"})
submission_ref = submitted.headers["X-Treg-Call-Id"]
await audit.drain()
async def fake_poll(*args, **kwargs):
return _response(status, {"id": "free-poll-task", "status": "processing"})
monkeypatch.setattr(call_service, "relay", fake_poll)
statements = []
def capture_sql(conn, cursor, statement, parameters, context, executemany):
statements.append(statement.lower())
event.listen(_engine.sync_engine, "before_cursor_execute", capture_sql)
try:
response = await clients.get("/call/replicate.predictions.get?id=free-poll-task")
await audit.drain()
finally:
event.remove(_engine.sync_engine, "before_cursor_execute", capture_sql)
money_sql = [sql for sql in statements
if re.search(r"\b(hold|ledgerentry|tagspend)\b", sql)
or sql.startswith("update org ")]
assert money_sql == []
assert response.status_code == status
assert response.headers["X-Treg-Cost-Micro"] == "0"
poll_ref = response.headers["X-Treg-Call-Id"]
async with session_maker() as db:
entries = (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == poll_ref))).scalars().all()
assert entries == []
assert await db.get(Hold, poll_ref) is None
assert await db.get(Hold, submission_ref) is not None
assert (await db.get(AsyncTaskRecord, submission_ref)).status == "pending"
@pytest.mark.parametrize("status", [200, 500])
async def test_owned_free_poll_hidden_before_activity_pagination(
clients: AsyncClient, monkeypatch, replicate_platform, status: int,
):
first = await _submit(clients, monkeypatch, {"id": "activity-task-1"})
await audit.drain()
second = await _submit(clients, monkeypatch, {"id": "activity-task-2"})
await audit.drain()
async def fake_poll(*args, **kwargs):
return _response(status, {"id": "activity-task-2", "status": "processing"})
monkeypatch.setattr(call_service, "relay", fake_poll)
polled = await clients.get("/call/replicate.predictions.get?id=activity-task-2")
assert polled.status_code == status
await audit.drain()
page = (await clients.get("/calls?limit=1&days=1")).json()
assert len(page) == 1 and page[0]["call_ref"] == second.headers["X-Treg-Call-Id"]
older = (await clients.get(f"/calls?limit=1&before_id={page[0]['id']}")).json()
assert len(older) == 1 and older[0]["call_ref"] == first.headers["X-Treg-Call-Id"]
async with session_maker() as db:
record = (await db.execute(select(CallRecord).where(
CallRecord.call_ref == polled.headers["X-Treg-Call-Id"]))).scalar_one()
assert record.kind == "async_poll"
assert record.cost_charged_micro == 0
if status == 500:
assert record.error_response
async def test_owned_free_poll_bypasses_reserve_and_reads_fresh_status(
clients: AsyncClient, monkeypatch, replicate_platform,
):
await _submit(clients, monkeypatch, {"id": "fresh-poll-task"})
async def must_not_reserve(*args, **kwargs):
raise AssertionError("free polling reached spend caps, holds or auto-top-up")
replies = iter(["processing", "succeeded"])
async def fake_poll(*args, **kwargs):
return _response(200, {"id": "fresh-poll-task", "status": next(replies)})
monkeypatch.setattr(call_service, "_platform_reserve", must_not_reserve)
monkeypatch.setattr(call_service, "relay", fake_poll)
for expected in ("processing", "succeeded"):
response = await clients.get("/call/replicate.predictions.get?id=fresh-poll-task",
headers={"Idempotency-Key": "same-poll"})
assert response.status_code == 200
assert response.json()["status"] == expected
assert response.headers["X-Treg-Cost-Micro"] == "0"
async def test_owned_free_poll_timeout_keeps_diagnostics_without_money_or_activity(
clients: AsyncClient, monkeypatch, replicate_platform,
):
submitted = await _submit(clients, monkeypatch, {"id": "timeout-poll-task"})
async def timeout(*args, **kwargs):
raise ReadTimeout("poll timed out")
monkeypatch.setattr(call_service, "relay", timeout)
response = await clients.get("/call/replicate.predictions.get?id=timeout-poll-task")
assert response.status_code == 502
assert response.headers["X-Treg-Cost-Micro"] == "0"
await audit.drain()
activity = (await clients.get("/calls")).json()
assert [row["call_ref"] for row in activity] == [submitted.headers["X-Treg-Call-Id"]]
async with session_maker() as db:
row = (await db.execute(select(CallRecord).where(
CallRecord.kind == "async_poll"))).scalar_one()
assert row.error_response and row.cost_charged_micro == 0
assert not (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == row.call_ref))).scalars().all()
@pytest.mark.parametrize("terminal, expected, kind, charged", [
("succeeded", "settled", "settle", 3000),
("failed", "released", "release", 0),
("canceled", "released", "release", 0),
])
async def test_owned_terminal_poll_finalizes_original_task_before_response(
clients: AsyncClient, monkeypatch, replicate_platform, terminal, expected, kind, charged,
):
submitted = await _submit(clients, monkeypatch, {"id": "instant-task"})
call_id = submitted.headers["X-Treg-Call-Id"]
document = {"id": "instant-task", "status": terminal,
"output": ["https://example.invalid/first.png"]}
async def fake_poll(*args, **kwargs):
return _response(200, document)
monkeypatch.setattr(call_service, "relay", fake_poll)
polled = await clients.get("/call/replicate.predictions.get?id=instant-task")
assert polled.status_code == 200 and polled.json() == document
assert polled.headers["X-Treg-Cost-Micro"] == "0"
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
assert row.status == expected and row.settled_micro == charged
assert row.completed_at is not None and row.attempts == 0
assert await db.get(Hold, call_id) is None
await audit.drain()
activity = (await clients.get("/calls")).json()
assert len(activity) == 1 and activity[0]["call_ref"] == call_id
assert activity[0]["async_task"]["status"] == expected
assert activity[0]["cost_charged_micro"] == charged
if terminal == "succeeded":
assert activity[0]["async_task"]["result_url"] == document["output"][0]
# Later observations must not charge again or replace the evidence used for settlement.
document = {**document, "output": ["https://example.invalid/later.png"]}
assert (await clients.get("/call/replicate.predictions.get?id=instant-task")).status_code == 200
async with session_maker() as db:
entries = (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == call_id, LedgerEntry.kind == kind))).scalars().all()
assert len(entries) == 1
archived = await archive.load_terminal_responses([(call_id, EP)])
assert json.loads(archived[call_id])["output"] == ["https://example.invalid/first.png"]
assert (await task_app.settle_due()).claimed == 0
async def test_firecrawl_crawl_poll_settles_reported_credits_once(clients: AsyncClient, monkeypatch):
monkeypatch.setenv("TREG_PLATFORM_KEY_FIRECRAWL", "PLATFORM-FIRECRAWL")
monkeypatch.setenv("TREG_PLATFORM_PROVIDERS", "firecrawl")
get_settings.cache_clear()
job_id = "01a0e894-a4fd-77fe-8131-e3f2d5720dbe"
try:
async def submit(*args, **kwargs):
return _response(200, {"success": True, "id": job_id,
"url": f"https://api.firecrawl.dev/v2/crawl/{job_id}"})
monkeypatch.setattr(call_service, "relay", submit)
org_id = (await clients.get("/orgs")).json()[0]["org_id"]
balance_url = f"/orgs/{org_id}/balance"
before = (await clients.get(balance_url)).json()["balance_micro"]
submitted = await clients.post("/call/firecrawl.web.crawl", json={
"url": "https://example.com", "limit": 3, "scrapeOptions": {"parsers": []},
})
assert submitted.status_code == 200, submitted.text
assert (await clients.get(balance_url)).json()["balance_micro"] == before - 15000
async def poll(*args, **kwargs):
return _response(200, {"success": True, "status": "completed", "total": 3,
"completed": 2, "creditsUsed": 2, "data": []})
monkeypatch.setattr(call_service, "relay", poll)
for _ in range(2):
response = await clients.get(f"/call/firecrawl.web.crawl.status?id={job_id}")
assert response.status_code == 200, response.text
assert response.json()["creditsUsed"] == 2
assert (await clients.get(balance_url)).json()["balance_micro"] == before - 10000
finally:
get_settings.cache_clear()
@pytest.mark.parametrize("status, content_type, body", [
(201, b"application/json", b"{}"),
(200, b"text/html", b"<html>WAF challenge</html>"),
])
async def test_a_2xx_without_a_readable_task_settles_at_zero_on_the_request_path(
clients: AsyncClient, monkeypatch, replicate_platform, status, content_type, body,
):
"""No task in the answer (none named, or not JSON at all) means nothing to poll and nothing to
charge: closed now, not parked until the 24-hour deadline (which is what an extraction failure
used to do)."""
async def fake_relay(*args, **kwargs):
async def stream():
yield body
async def close():
return None
return UpstreamResponse(status, ((b"content-type", content_type),), stream(), close)
monkeypatch.setattr(call_service, "relay", fake_relay)
response = await clients.post(f"/call/{EP}", json={"input": {"prompt": "x", "num_outputs": 1}})
assert response.status_code == status
assert response.headers["X-Treg-Cost-Micro"] == "0"
call_id = response.headers["X-Treg-Call-Id"]
async with session_maker() as db:
assert await db.get(AsyncTaskRecord, call_id) is None
assert await db.get(Hold, call_id) is None
entries = {e.kind: e.amount_micro for e in (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == call_id))).scalars().all()}
assert entries == {"reserve": -3000, "settle": 0}
async def test_one_failing_row_does_not_abort_the_tick(
clients: AsyncClient, monkeypatch, replicate_platform,
):
"""relay() raises GatewayFailed (an unset platform key, an SSRF refusal), which is not an
httpx or JSON error: it must back that row off and let the tick serve the others."""
from treg.application.call.types import GatewayFailed
broken = await _due_submission(clients, monkeypatch, {"status": "succeeded", "output": ["u"]})
fine = await _due_submission(clients, monkeypatch, {"status": "succeeded", "output": ["u"]})
async def poll(row, client):
if row.call_id == broken:
raise GatewayFailed("injection_failed", status_code=502, detail="no platform key")
return 200, json.dumps({"status": "succeeded", "output": ["u"]}).encode()
monkeypatch.setattr(task_app, "_poll", poll)
result = await task_app.settle_due()
assert (result.claimed, result.settled, result.backed_off) == (2, 1, 1)
async with session_maker() as db:
assert (await db.get(AsyncTaskRecord, broken)).status == "pending"
assert (await db.get(AsyncTaskRecord, fine)).status == "settled"
async def test_pending_row_write_failure_releases_the_hold_and_alerts(
clients: AsyncClient, monkeypatch, replicate_platform,
):
async def fail_persistence(*args, **kwargs):
raise RuntimeError("database unavailable")
monkeypatch.setattr(task_app, "defer_submission", fail_persistence)
response = await _submit(clients, monkeypatch, {
"id": "prediction-untracked",
"urls": {"get": "https://api.replicate.com/v1/predictions/untracked"},
})
assert response.status_code == 201
assert response.headers["X-Treg-Cost-Micro"] == "0"
call_id = response.headers["X-Treg-Call-Id"]
async with session_maker() as db:
assert await db.get(AsyncTaskRecord, call_id) is None
assert await db.get(Hold, call_id) is None
entry = (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == call_id, LedgerEntry.kind == "release"))).scalar_one()
# treg's own failure is treg's cost: the whole reserve goes back, nothing is settled.
assert entry.amount_micro == 3000
assert entry.meta.get("reason") == "async_task_not_recorded"
@pytest.fixture
def minimax_platform(monkeypatch):
monkeypatch.setenv("TREG_PLATFORM_KEY_MINIMAX", "test-platform-token")
monkeypatch.setenv("TREG_PLATFORM_PROVIDERS", "minimax")
get_settings.cache_clear()
yield
get_settings.cache_clear()
@pytest.fixture
def openrouter_platform(monkeypatch):
monkeypatch.setenv("TREG_PLATFORM_KEY_OPENROUTER", "test-platform-token")
monkeypatch.setenv("TREG_PLATFORM_PROVIDERS", "openrouter")
get_settings.cache_clear()
yield
get_settings.cache_clear()
@pytest.fixture
def legacy_async_platform(monkeypatch):
for provider in ("apify", "brightdata", "companyenrich", "oceanio"):
monkeypatch.setenv(f"TREG_PLATFORM_KEY_{provider.upper()}", "test-platform-token")
monkeypatch.setenv("TREG_PLATFORM_PROVIDERS", "apify,brightdata,companyenrich,oceanio")
get_settings.cache_clear()
yield
get_settings.cache_clear()
async def test_platform_openrouter_model_is_bound_to_the_selected_catalog_row(
clients: AsyncClient, monkeypatch, openrouter_platform,
):
relayed = []
async def fake_relay(*args, **kwargs):
relayed.append(True)
return _response(201, {"id": "video-owned"})
monkeypatch.setattr(call_service, "relay", fake_relay)
endpoint = "/call/openrouter.x.google-veo-3-1-lite"
body = {"model": "black-forest-labs/flux-3-video", "prompt": "A paper boat.",
"duration": 4, "resolution": "720p", "generate_audio": True}
refused = await clients.post(endpoint, json=body)
assert refused.status_code == 400
assert refused.json()["detail"]["parameter"] == "body.model"
assert relayed == []
async with session_maker() as db:
assert (await db.execute(select(Hold))).scalars().all() == []
body["model"] = "google/veo-3.1-lite"
accepted = await clients.post(endpoint, json=body)
assert accepted.status_code == 201 and relayed == [True]
async with session_maker() as db:
row = (await db.execute(select(AsyncTaskRecord).where(
AsyncTaskRecord.task_id == "video-owned"))).scalar_one()
assert row.endpoint_id == "openrouter.x.google-veo-3-1-lite"
async def test_platform_model_selector_rejects_duplicate_json_keys(
clients: AsyncClient, monkeypatch, openrouter_platform,
):
async def must_not_relay(*args, **kwargs):
raise AssertionError("ambiguous model reached the provider")
monkeypatch.setattr(call_service, "relay", must_not_relay)
body = (b'{"model":"google/veo-3.1-lite",'
b'"model":"black-forest-labs/flux-3-video","prompt":"x"}')
response = await clients.post(
"/call/openrouter.x.google-veo-3-1-lite", content=body,
headers={"content-type": "application/json"})
assert response.status_code == 400
assert "repeats JSON field" in response.json()["detail"]
async def test_byok_openrouter_remains_a_faithful_relay_for_model_choice(
clients: AsyncClient, monkeypatch, openrouter_platform,
):
await clients.post("/secrets", json={"name": "openrouter", "value": "own-token"})
relayed = []
async def fake_relay(*args, **kwargs):
relayed.append(True)
return _response(201, {"id": "byok-video"})
monkeypatch.setattr(call_service, "relay", fake_relay)
response = await clients.post("/call/openrouter.x.google-veo-3-1-lite", json={
"model": "black-forest-labs/flux-3-video", "prompt": "A paper boat."})
assert response.status_code == 201 and relayed == [True]
async with session_maker() as db:
assert (await db.execute(select(AsyncTaskRecord))).scalars().all() == []
async def test_platform_task_status_requires_same_org_submission(
clients: AsyncClient, monkeypatch, replicate_platform,
):
submitted = await _submit(clients, monkeypatch, {
"id": "prediction-owned",
"urls": {"get": "https://api.replicate.com/v1/predictions/owned"},
})
assert submitted.status_code == 201
relayed = []
async def fake_status(*args, **kwargs):
relayed.append(True)
return _response(200, {"id": "prediction-owned", "status": "processing"})
monkeypatch.setattr(call_service, "relay", fake_status)
own = await clients.get("/call/replicate.predictions.get?id=prediction-owned")
assert own.status_code == 200 and relayed == [True]
unknown = await clients.get("/call/replicate.predictions.get?id=prediction-unknown")
assert unknown.status_code == 403 and relayed == [True]
assert unknown.json()["detail"]["error"] == "async_resource_not_owned"
ambiguous = await clients.get(
"/call/replicate.predictions.get",
params=[("id", "prediction-owned"), ("id", "prediction-unknown")])
assert ambiguous.status_code == 400 and relayed == [True]
other = await clients.post("/users", json={"email": "task-stranger@example.com"})
stranger = {"X-Treg-Token": other.json()["token"]}
denied = await clients.get(
"/call/replicate.predictions.get?id=prediction-owned", headers=stranger)
assert denied.status_code == 403 and relayed == [True]
assert denied.json()["detail"] == unknown.json()["detail"]
@pytest.mark.parametrize("platform, secret, url", [
("replicate_platform", "replicate",
"/call/replicate.predictions.get?id=arbitrary-own-account-id"),
("legacy_async_platform", "apify", "/call/apify.web.scrape.job.status?run_id=arbitrary"),
])
async def test_byok_task_status_keeps_direct_provider_object_access(
clients: AsyncClient, monkeypatch, request, platform, secret, url,
):
"""A team's own key reaches any object on its own provider account; ownership checks guard
only treg's shared key."""
request.getfixturevalue(platform)
await clients.post("/secrets", json={"name": secret, "value": "own-token"})
async def fake_status(*args, **kwargs):
return _response(200, {"id": "arbitrary-own-account-id", "status": "processing"})
monkeypatch.setattr(call_service, "relay", fake_status)
assert (await clients.get(url)).status_code == 200
@pytest.mark.parametrize(("start", "payload", "created", "owned_calls"), [
(
"/call/brightdata.web.scrape.job.start?dataset_id=gd_test", [{"url": "https://example.com"}],
{"snapshot_id": "snapshot-owned"},
[
"/call/brightdata.web.scrape.job.status?snapshot_id=snapshot-owned",
"/call/brightdata.web.scrape.job.results?snapshot_id=snapshot-owned&format=json",
],
),
(
"/call/companyenrich.companies.enrich.bulk.start", {"domains": ["example.com"]},
{"job_id": "job-owned", "status": "pending"},
["/call/companyenrich.companies.enrich.bulk.status?jobId=job-owned"],
),
(
"/call/companyenrich.companies.search.async.start",
{"count": 1, "search": {"countries": ["US"]}},
{"job_id": "company-search-owned", "status": "pending"},
["/call/companyenrich.companies.search.async.status?jobId=company-search-owned"],
),
(
"/call/companyenrich.people.email.bulk.start",
{"items": [{"person_id": 1, "domain": "example.com"}]},
{"job_id": "people-email-owned", "status": "pending"},
["/call/companyenrich.people.email.bulk.status?jobId=people-email-owned"],
),
(
"/call/companyenrich.people.search.async.start", {"count": 1, "domains": ["example.com"]},
{"job_id": "people-search-owned", "status": "pending"},
["/call/companyenrich.people.search.async.status?jobId=people-search-owned"],
),
(
"/call/oceanio.companies.segment.create", {"domains": ["example.com"]},
{"segmentationId": 12345},
["/call/oceanio.companies.segment.get?segmentation_id=12345"],
),
])
async def test_legacy_platform_async_resources_are_recorded_and_authorized(
clients: AsyncClient, monkeypatch, legacy_async_platform,
start: str, payload: object, created: dict, owned_calls: list[str],
):
responses = [created, {"status": "running"}, []]
async def fake_relay(*args, **kwargs):
return _response(200, responses.pop(0))
monkeypatch.setattr(call_service, "relay", fake_relay)
submitted = await clients.post(start, json=payload)
assert submitted.status_code == 200
for url in owned_calls:
assert (await clients.get(url)).status_code == 200
other = await clients.post("/users", json={"email": "legacy-stranger@example.com"})
denied = await clients.get(owned_calls[0], headers={"X-Treg-Token": other.json()["token"]})
assert denied.status_code == 403
async with session_maker() as db:
records = (await db.execute(select(AsyncResourceRecord))).scalars().all()
assert records
@pytest.mark.parametrize("url", [
"/call/apify.web.scrape.job.status?run_id=unknown",
"/call/oceanio.companies.segment.get?segmentation_id=99999",
])
async def test_legacy_platform_async_utilities_deny_unknown_ids_before_relay(
clients: AsyncClient, monkeypatch, legacy_async_platform, url: str,
):
async def must_not_relay(*args, **kwargs):
raise AssertionError("unowned async resource reached the shared provider account")
monkeypatch.setattr(call_service, "relay", must_not_relay)
response = await clients.get(url)
assert response.status_code == 403
assert response.json()["detail"]["error"] == "async_resource_not_owned"
async def test_icypeas_shared_key_reads_only_this_teams_searches(clients: AsyncClient, monkeypatch):
"""The listing modes enumerate every search on treg's one Icypeas account: every team's."""
monkeypatch.setenv("TREG_PLATFORM_KEY_ICYPEAS", "test-platform-token")
monkeypatch.setenv("TREG_PLATFORM_PROVIDERS", "icypeas")
get_settings.cache_clear()
responses = [{"success": True, "item": {"_id": "search-owned", "status": "NONE"}},
{"success": True, "status": "in_progress", "file": "file-owned"}]
async def fake_relay(*args, **kwargs):
return _response(200, responses.pop(0) if responses else {"success": True, "items": []})
monkeypatch.setattr(call_service, "relay", fake_relay)
assert (await clients.post("/call/icypeas.people.email.find", json={
"firstname": "A", "lastname": "B", "domainOrCompany": "example.com"})).status_code == 200
assert (await clients.post("/call/icypeas.bulk.search", json={
"name": "x", "task": "email-verification", "data": [["a@example.com"]]})).status_code == 200
owned = [("icypeas.search.results.read", {"id": "search-owned"}),
("icypeas.bulk.results.read", {"mode": "bulk", "file": "file-owned"}),
("icypeas.search.files.read", {"file": "file-owned"})]
for endpoint, body in owned:
assert (await clients.post(f"/call/{endpoint}", json=body)).status_code == 200, endpoint
for endpoint, body in [("icypeas.search.results.read", {"mode": "single", "type": "email-search"}),
("icypeas.search.results.read", {"id": "search-owned", "mode": "single"}),
("icypeas.bulk.results.read", {"mode": "bulk"}),
("icypeas.search.files.read", {}),
("icypeas.search.results.read", {"id": "someone-elses"})]:
response = await clients.post(f"/call/{endpoint}", json=body)
assert response.status_code in (400, 403), (endpoint, body, response.status_code)
other = await clients.post("/users", json={"email": "icypeas-stranger@example.com"})
stranger = {"X-Treg-Token": other.json()["token"]}
for endpoint, body in owned:
assert (await clients.post(f"/call/{endpoint}", json=body, headers=stranger)).status_code == 403
get_settings.cache_clear()
async def test_apify_actor_start_needs_own_key(
clients: AsyncClient, monkeypatch, legacy_async_platform,
):
async def must_not_relay(*args, **kwargs):
raise AssertionError("an unmetered actor run reached treg's Apify account")
monkeypatch.setattr(call_service, "relay", must_not_relay)
response = await clients.post(
"/call/apify.web.scrape.job.start?actor_id=apify~hello-world", json={})
assert response.status_code == 404
async def test_legacy_platform_async_mutation_denies_unknown_resource_before_relay(
clients: AsyncClient, monkeypatch, legacy_async_platform,
):
async def must_not_relay(*args, **kwargs):
raise AssertionError("unowned segmentation reached the shared provider account")
monkeypatch.setattr(call_service, "relay", must_not_relay)
response = await clients.post(
"/call/oceanio.companies.segment.mark_domains?segmentation_id=99999",
json={"domains": ["example.com"], "type": "positive"},
)
assert response.status_code == 403
async def _submit_minimax(clients: AsyncClient, monkeypatch, task_id: str) -> str:
async def fake_submit(*args, **kwargs):
return _response(200, {"task_id": task_id, "base_resp": {"status_code": 0}})
monkeypatch.setattr(call_service, "relay", fake_submit)
response = await clients.post("/call/minimax.video-gen.from_text", json={
"model": "MiniMax-Hailuo-2.3", "prompt": "A paper boat.", "duration": 6,
"resolution": "768P"})
assert response.status_code == 200
return response.headers["X-Treg-Call-Id"]
async def test_owned_terminal_poll_teaches_result_id_before_fetch(
clients: AsyncClient, monkeypatch, minimax_platform,
):
call_id = await _submit_minimax(clients, monkeypatch, "minimax-task-owned")
relay_calls = []
async def fake_relay(*args, **kwargs):
relay_calls.append(True)
if len(relay_calls) == 1:
return _response(200, {"status": "Success", "file_id": "minimax-file-owned"})
return _response(200, {"file": {"download_url": "https://example.invalid/video.mp4"}})
monkeypatch.setattr(call_service, "relay", fake_relay)
polled = await clients.get(
"/call/minimax.video-gen.task.status?task_id=minimax-task-owned")
assert polled.status_code == 200
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
assert row.result_id == "minimax-file-owned"
fetched = await clients.get(
"/call/minimax.video-gen.result.retrieve?file_id=minimax-file-owned")
assert fetched.status_code == 200 and len(relay_calls) == 2
other = await clients.post("/users", json={"email": "file-stranger@example.com"})
denied = await clients.get(
"/call/minimax.video-gen.result.retrieve?file_id=minimax-file-owned",
headers={"X-Treg-Token": other.json()["token"]})
assert denied.status_code == 403 and len(relay_calls) == 2
async def test_worker_terminal_success_persists_fetch_result_ownership(
clients: AsyncClient, monkeypatch, minimax_platform,
):
call_id = await _submit_minimax(clients, monkeypatch, "minimax-task-worker")
outcome = await task_app._finish(
call_id, "success", {"status": "Success", "file_id": "minimax-file-worker"},
utcnow_naive())
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
assert outcome == "settled" and row.result_id == "minimax-file-worker"
async def test_a_2xx_that_fails_the_expect_rule_releases_and_is_not_deferred(
clients: AsyncClient, monkeypatch, minimax_platform,
):
"""MiniMax answers HTTP 200 with the error in the envelope (live 2026-09-02: base_resp 2013,
"model MiniMax-Hailuo-2.3-Fast does not support Text-to-Video mode"). No task exists, so
nothing may be deferred and nothing may be charged."""
async def fake_relay(*args, **kwargs):
return _response(200, {"task_id": "", "base_resp": {
"status_code": 2013, "status_msg": "invalid params"}})
monkeypatch.setattr(call_service, "relay", fake_relay)
response = await clients.post("/call/minimax.video-gen.from_text", json={
"model": "MiniMax-Hailuo-2.3", "prompt": "A paper boat.", "duration": 6,
"resolution": "768P"})
assert response.status_code == 200
assert response.headers["X-Treg-Cost-Micro"] == "0"
call_id = response.headers["X-Treg-Call-Id"]
async with session_maker() as db:
assert await db.get(AsyncTaskRecord, call_id) is None
assert await db.get(Hold, call_id) is None
entries = {e.kind: e.amount_micro for e in (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == call_id))).scalars().all()}
# A failed envelope is a per_success miss: settled at zero, the whole reserve given back.
assert entries == {"reserve": -280000, "settle": 0}
async def _due_submission(clients, monkeypatch, document: dict) -> str:
response = await _submit(clients, monkeypatch, {
"id": "prediction-worker",
"urls": {"get": "https://api.replicate.com/v1/predictions/worker"},
})
call_id = response.headers["X-Treg-Call-Id"]
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
row.next_check_at = utcnow_naive() - timedelta(seconds=1)
await db.commit()
async def fake_poll(row, client):
return 200, json.dumps(document).encode()
monkeypatch.setattr(task_app, "_poll", fake_poll)
return call_id
async def test_worker_settles_terminal_success(
clients: AsyncClient, monkeypatch, replicate_platform,
):
call_id = await _due_submission(clients, monkeypatch, {"status": "succeeded", "output": ["url"]})
result = await task_app.settle_due()
assert result.settled == 1
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
hold = await db.get(Hold, call_id)
entry = (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == call_id, LedgerEntry.kind == "settle"))).scalar_one()
assert row.status == "settled" and row.settled_micro == 3000
assert hold is None and entry.amount_micro == -3000
async with session_maker() as db:
key = (await db.execute(select(ArchiveKey).where(
ArchiveKey.req_url == f"treg://asynctasks/{call_id}"))).scalar_one()
snapshot = (await db.execute(select(ArchiveSnapshot).where(
ArchiveSnapshot.key_id == key.id))).scalar_one()
report = await reconcile.async_task_settlement(
db, utcnow_naive() - timedelta(hours=1))
assert json.loads(snapshot.body)["status"] == "succeeded"
assert report["providers"] == [{
"provider": "replicate", "successes": 1, "failures": 0,
"settled_micro": 3000, "tasks": 1, "success_rate": 1.0,
"settled_usd": 0.003,
}]
async def test_worker_releases_terminal_failure(
clients: AsyncClient, monkeypatch, replicate_platform,
):
call_id = await _due_submission(clients, monkeypatch, {"status": "failed", "error": "rejected"})
result = await task_app.settle_due()
assert result.released == 1
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
hold = await db.get(Hold, call_id)
assert row.status == "released" and row.settled_micro == 0 and hold is None
async def test_worker_backs_off_unknown_status(
clients: AsyncClient, monkeypatch, replicate_platform,
):
call_id = await _due_submission(clients, monkeypatch, {"status": "provider_added_a_state"})
before = utcnow_naive()
result = await task_app.settle_due()
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
hold = await db.get(Hold, call_id)
assert result.backed_off == 1
assert row.status == "pending" and row.next_check_at > before and hold is not None
async with session_maker() as db:
hold = await db.get(Hold, call_id)
hold.created_at = utcnow_naive() - timedelta(seconds=ledger.hold_ttl_s() + 1)
await db.commit()
assert await ledger.reap_stale_holds(db, org_id=row.org_id) == 0
assert await db.get(Hold, call_id) is not None
async def test_worker_timeout_releases_the_hold_and_flags_it_for_review(
clients: AsyncClient, monkeypatch, replicate_platform, caplog,
):
"""An outcome nobody observed is the platform's cost, never the customer's."""
call_id = await _due_submission(clients, monkeypatch, {"status": "processing"})
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
row.created_at = utcnow_naive() - asynctasks.MAX_AGE - timedelta(seconds=1)
row.next_check_at = utcnow_naive() - timedelta(seconds=1)
await db.commit()
with caplog.at_level("ERROR", logger="treg.asynctasks"):
result = await task_app.settle_due()
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
hold = await db.get(Hold, call_id)
entry = (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == call_id, LedgerEntry.kind == "release"))).scalar_one()
assert result.timed_out == 1
assert row.status == "timed_out" and row.settled_micro == 0 and row.reserved_micro == 3000
assert hold is None and entry.meta.get("reconcile_review") is True
assert any("ASYNC TASK TIMED OUT" in rec.message for rec in caplog.records)
async with session_maker() as db:
report = await reconcile.async_task_settlement(
db, utcnow_naive() - timedelta(hours=1))
assert [item["call_id"] for item in report["absorbed_timeouts"]] == [call_id]
assert report["absorbed_timeouts"][0]["reserved_micro"] == 3000
def test_usage_terms_sum_every_reported_meter_at_its_rate():
"""A response that reports several token meters (Gemini's usageMetadata) settles on the sum
of each meter times its rate; per-modality entries are selected by key, never by index."""
cost = {"settle": "usage", "fallback": {"value": 0.15}, "usage": {"unit": "usd", "terms": [
{"path": "usageMetadata.promptTokenCount", "rate": 0.000002},
{"path": "usageMetadata.candidatesTokenCount", "rate": 0.000012},
{"path": "usageMetadata.thoughtsTokenCount", "rate": 0.000012},
{"path": "usageMetadata.candidatesTokensDetails[modality=IMAGE].tokenCount",
"rate": 0.000108}]}}
basis = settlement.derive_basis(
cost, request={}, input_schema={}, unit_micro=1_000_000, terminal=False)
usage = {"promptTokenCount": 17, "candidatesTokenCount": 1229, "thoughtsTokenCount": 141,
"candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 109},
{"modality": "IMAGE", "tokenCount": 1120}]}
# 17*2 + (1229+141)*12 + 1120*108 = 137,434 micro-USD: the image meter at $120/M in total.
assert settlement.settle(basis, {"terminal": {"usageMetadata": usage}}) == 137_434
# proto3 JSON omits a zero meter: a blocked prompt reports only its input and pays for it.
assert settlement.settle(basis, {"terminal": {"usageMetadata": {"promptTokenCount": 40}}}) == 80
# No meter at all is unobserved, never free: the success settles at the reserve.
assert settlement.usage_evidence(basis, {"terminal": {"candidates": []}}) is None
assert settlement.settle(basis, {"terminal": None}) == basis["reserve_micro"] == 150_000
# A malformed meter poisons the figure rather than silently billing less.
bad = {"usageMetadata": {**usage, "thoughtsTokenCount": -1}}
assert settlement.usage_evidence(basis, {"terminal": bad}) is None
# A spooled answer keeps only the keys these meters start at.
assert settlement.usage_roots(cost) == ("usageMetadata",)
assert settlement.usage_roots({"usage": {"path": "usage.cost", "unit": "usd"}}) == ("usage",)
def test_basis_derivation_and_settlement_table_vs_usage():
table_cost = {"table": [{"when": {"body.n": 2}, "value": 0.01}],
"fallback": {"value": 0.04}, "settle": "table"}
table = settlement.derive_basis(
table_cost, request={"body": {"n": 2}}, input_schema={}, unit_micro=1_000_000,
terminal=True)
assert table["amount"]["kind"] == "table"
assert settlement.settle(table, {"terminal": {}}) == 10_000
usage_cost = {"settle": "usage", "usage": {"path": "usage.cost", "unit": "usd"},
"fallback": {"value": 1.0}}
usage = settlement.derive_basis(
usage_cost, request={}, input_schema={}, unit_micro=1_000_000, terminal=True)
assert usage["amount"]["kind"] == "usage"
assert settlement.settle(usage, {"terminal": {"usage": {"cost": 0.125}}}) == 125_000
# No table row matched and no usage figure: the reserve is the fallback, and a success with
# no usage evidence settles at that reserve, never above it.
assert usage["reserve_micro"] == usage["fallback_micro"] == 1_000_000
assert settlement.usage_evidence(usage, {"terminal": {"status": "completed"}}) is None
assert settlement.settle(usage, {"terminal": {"status": "completed"}}) == 1_000_000
# A usage row reserves what the rate card says THIS request costs, not the matrix ceiling.
rate_card = {"settle": "usage", "usage": {"path": "usage.cost", "unit": "usd"},
"table": [{"when": {"body.resolution": "480p"}, "value": 0.05, "times": "body.duration"},
{"when": {"body.resolution": "1080p"}, "value": 0.20, "times": "body.duration"}],
"fallback": {"value": 6.0}}
schema = {"body": {"resolution": {"type": "string"}, "duration": {"type": "integer", "max": 30}}}
cheap = settlement.derive_basis(
rate_card, request={"body": {"resolution": "480p", "duration": 2}}, input_schema=schema,
unit_micro=1_000_000, terminal=True)
assert cheap["reserve_micro"] == 100_000 and cheap["fallback_micro"] == 6_000_000
# The provider's reported cost settles even when it exceeds the reserve (Wan 3.0's minimum).
assert settlement.settle(cheap, {"terminal": {"usage": {"cost": 0.2125}}}) == 212_500
# A provider that meters in its own credits settles the reported credits at the rate frozen
# into the basis; a sentinel the provider demands (`duration: -1`, auto) is priced by its own
# row, so it neither multiplies the rate negative nor turns the ceiling into the bill.
credits = {"settle": "usage", "usage": {"path": "usage.credits", "unit": "credit"},
"table": [{"when": {"body.resolution": "480p", "body.duration": -1}, "value": 3.558},
{"when": {"body.resolution": "480p"}, "value": 0.1186,
"times": "body.duration", "times_min": 4}],
"fallback": {"value": 13.87}}
schema = {"body": {"resolution": {"type": "string"},
"duration": {"type": "integer", "min": -1, "max": 30}}}
auto = settlement.derive_basis(
credits, request={"body": {"resolution": "480p", "duration": -1}}, input_schema=schema,
unit_micro=1_000_000, terminal=True, usage_unit_micro=1_000)
assert auto["reserve_micro"] == 3_558_000
assert settlement.settle(auto, {"terminal": {"usage": {"credits": 712}}}) == 712_000
# A declared minimum of -1 never lets zero multiply a rate: the ceiling is held instead.
zero = settlement.derive_basis(
credits, request={"body": {"resolution": "480p", "duration": 0}}, input_schema=schema,
unit_micro=1_000_000, terminal=True, usage_unit_micro=1_000)
assert zero["reserve_micro"] == 13_870_000
# A credit meter with no frozen rate cannot be priced: the reserve settles, not credits-as-USD.
unrated = settlement.derive_basis(
credits, request={"body": {"resolution": "480p", "duration": 5}}, input_schema=schema,
unit_micro=1_000_000, terminal=True)
assert settlement.settle(unrated, {"terminal": {"usage": {"credits": 712}}}) == 593_000
request = settlement.request_evidence(
[("id", "42"), ("count", "2")], b"{}", path_names={"id"})
path_table = {"table": [{"when": {"pathParams.id": 42}, "value": 0.01,
"times": "queryParams.count"}],
"fallback": {"value": 0.10}, "settle": "table"}
schema = {"pathParams": {"id": {"type": "integer"}},
"queryParams": {"count": {"type": "integer"}}}
basis = settlement.derive_basis(
path_table, request=request, input_schema=schema, unit_micro=1_000_000, terminal=True)
assert basis["reserve_micro"] == 20_000
def test_terminal_classification_coerces_status_values_and_treats_none_as_progress():
descriptor = {
"status": {
"path": "task.status",
"success": [2],
"failure": ["3"],
"billed_failure": [4],
},
}
assert asynctasks.classify_terminal(descriptor, {"task": {"status": "2"}}) == "success"
assert asynctasks.classify_terminal(descriptor, {"task": {"status": 3}}) == "failure"
assert asynctasks.classify_terminal(descriptor, {"task": {"status": "4"}}) == "billed_failure"
assert asynctasks.classify_terminal(descriptor, {"task": {"status": None}}) == "progress"
assert asynctasks.classify_terminal(descriptor, {"task": {}}) == "progress"
async def _activity_row(clients: AsyncClient, call_id: str) -> dict:
await audit.drain()
await archive.drain()
rows = (await clients.get("/calls")).json()
return next(row for row in rows if row["call_ref"] == call_id)
async def test_activity_reports_task_state_and_artifact(
clients: AsyncClient, monkeypatch, replicate_platform,
):
"""The audit row froze the reserve as the charge; the feed must show what actually happened."""
call_id = await _due_submission(clients, monkeypatch, {
"status": "succeeded", "output": ["https://replicate.delivery/out.webp"],
# Keep the terminal envelope above archive._COMPRESS_MIN_BYTES. Production video status
# bodies are compressed, which is where the activity reader once returned encoded bytes
# to json.loads and silently lost the artifact.
"provider_metadata": "x" * 512,
})
pending = await _activity_row(clients, call_id)
assert pending["cost_charged_micro"] is None
assert pending["async_task"]["status"] == "pending"
assert pending["async_task"]["reserved_micro"] == 3000
assert pending["async_task"]["result_url"] is None
assert (await task_app.settle_due()).settled == 1
settled = await _activity_row(clients, call_id)
assert settled["cost_charged_micro"] == 3000
task = settled["async_task"]
assert task["status"] == "settled" and task["settled_micro"] == 3000
assert task["result_url"] == "https://replicate.delivery/out.webp"
assert task["completed_at"] is not None
one = (await clients.get(f"/calls/{call_id}")).json()
assert one["async_task"]["result_url"] == "https://replicate.delivery/out.webp"
assert one["call"]["cost_charged_micro"] == 3000 and one["charged_micro"] == 3000
def test_artifact_reads_both_result_modes():
by_path = {"result": {"path": "task.content.url", "ttl_note": "time-limited"}}
found = asynctasks.artifact(by_path, {"task": {"content": {"url": "https://x.invalid/v.mp4"}}})
assert found["result_url"] == "https://x.invalid/v.mp4" and found["ttl_note"] == "time-limited"
assert found["fetch"] is None
by_fetch = {"result": {"fetch": "minimax.video-gen.result.retrieve",
"fetch_param": {"in": "queryParams", "name": "file_id",
"value_from": "file_id"},
"ttl_note": "9h"}}
found = asynctasks.artifact(by_fetch, {"status": "Success", "file_id": "f-1"})
assert found["result_url"] is None
assert found["fetch"] == {"endpoint": "minimax.video-gen.result.retrieve",
"name": "file_id", "value": "f-1"}
assert asynctasks.artifact(by_fetch, {"status": "Success"})["fetch"] is None
assert asynctasks.artifact({}, {"anything": 1})["result_url"] is None
async def test_query_parameter_poll_travels_as_query_items(monkeypatch):
"""MiniMax v1 polls `GET /v1/query/video_generation?task_id=…`. The relay builds the upstream
query from `query_items` only, so the id must ride there (live 2026-09-02: appended to the URL it
arrived empty and the provider answered 2013 "invalid params" on every tick)."""
seen = {}
async def fake_relay(request, url, tool, *args, **kwargs):
seen["url"], seen["query"] = url, request.query_items
async def stream():
yield b'{"status": "Success"}'
async def close():
return None
return UpstreamResponse(200, (), stream(), close)
monkeypatch.setattr(task_app, "relay", fake_relay)
row = AsyncTaskRecord(
call_id="q-1", org_id=1, provider="minimax", endpoint_id="minimax.video-gen.from_text",
task_id="437372532953204", reserved_micro=1, next_check_at=utcnow_naive(),
descriptor={"poll": {"endpoint": "minimax.video-gen.task.status",
"param": {"in": "queryParams", "name": "task_id"}}})
status, body = await task_app._poll(row, None)
assert status == 200 and body == b'{"status": "Success"}'
assert seen["url"] == "https://api.minimax.io/v1/query/video_generation"
assert seen["query"] == (("task_id", "437372532953204"),)
async def test_reconcile_lists_usage_overruns_and_platform_absorbed_shortfalls(clients: AsyncClient):
"""The two places a usage-settled task can cost more than its reserve, made visible: the team
paid the overrun from its balance; the platform absorbed whatever its blocks could not cover."""
now = utcnow_naive()
async with session_maker() as db:
db.add(AsyncTaskRecord(
call_id="over-1", org_id=1, provider="openrouter",
endpoint_id="openrouter.video-gen.wan-3-0.from_text", task_id="t", reserved_micro=100_000,
settled_micro=212_500, status="settled", created_at=now, next_check_at=now,
completed_at=now, descriptor={}, settlement_basis={}))
db.add(AsyncTaskRecord(
call_id="even-1", org_id=1, provider="openrouter",
endpoint_id="openrouter.x.google-veo-3-1", task_id="u", reserved_micro=800_000,
settled_micro=800_000, status="settled", created_at=now, next_check_at=now,
completed_at=now, descriptor={}, settlement_basis={}))
db.add(LedgerEntry(
id="le-over-1", org_id=1, kind="settle", amount_micro=-120_000, call_id="over-1",
endpoint_id="openrouter.video-gen.wan-3-0.from_text", created_at=now,
meta={"settled_micro": 212_500, "consumed_micro": 120_000, "block_shortfall_micro": 92_500}))
db.add(LedgerEntry(
id="le-even-1", org_id=1, kind="settle", amount_micro=-800_000, call_id="even-1",
endpoint_id="openrouter.x.google-veo-3-1", created_at=now,
meta={"settled_micro": 800_000, "consumed_micro": 800_000, "block_shortfall_micro": 0}))
await db.commit()
report = await reconcile.async_task_settlement(db, now - timedelta(hours=1))
assert [o["call_id"] for o in report["overruns"]] == ["over-1"]
assert report["overruns"][0]["overrun_micro"] == 112_500 and report["overruns"][0]["ratio"] == 2.125
assert report["overruns_by_endpoint"] == [{
"endpoint_id": "openrouter.video-gen.wan-3-0.from_text", "provider": "openrouter",
"tasks": 1, "overrun_micro": 112_500, "max_ratio": 2.125}]
assert [s["call_id"] for s in report["absorbed_shortfalls"]] == ["over-1"]
assert report["absorbed_shortfall_micro"] == 92_500
def test_times_multiplier_is_bounded_by_the_input_schema():
"""A caller cannot reserve zero with duration 0 or bill past the ceiling with duration 100:
an out-of-range, non-finite or non-positive multiplier matches no row and prices at the fallback."""
cost = {"table": [{"when": {"body.input.resolution": "480p"}, "value": 0.05,
"times": "body.input.duration"}],
"fallback": {"value": 6.0}}
schema = {"body": {"input": {"type": "object", "properties": {
"resolution": {"type": "string"},
"duration": {"type": "integer", "min": 2, "max": 30}}}}}
def price(duration):
return settlement.table_amount_micro(
cost, {"body": {"input": {"resolution": "480p", "duration": duration}}}, schema, 1_000_000)
assert price(2) == 100_000 and price(30) == 1_500_000
assert price(0) == price(-3) == price(31) == price(100) == price(float("nan")) == 6_000_000
assert price(float("inf")) == 6_000_000 and price("5") == 250_000 # strings coerce by type
async def test_worker_ignores_terminal_looking_fields_on_error_responses(
clients: AsyncClient, monkeypatch, replicate_platform,
):
call_id = await _due_submission(clients, monkeypatch, {"status": "succeeded", "output": ["u"]})
async def fake_poll(row, client):
return 404, json.dumps({"status": "succeeded", "output": ["u"]}).encode()
monkeypatch.setattr(task_app, "_poll", fake_poll)
result = await task_app.settle_due()
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
hold = await db.get(Hold, call_id)
assert result.backed_off == 1 and row.status == "pending" and hold is not None
def test_poll_target_follows_the_declared_location_and_encodes_path_values():
row = AsyncTaskRecord(
call_id="p-1", org_id=1, provider="minimax", endpoint_id="minimax.video-gen.h3.generate",
task_id="a?view=admin/../x", reserved_micro=1, next_check_at=utcnow_naive(),
descriptor={"poll": {"endpoint": "minimax.video-gen.v2.task.status",
"param": {"in": "pathParams", "name": "task_id"}}})
method, url, query = task_app._poll_target(row)
assert (method, query) == ("GET", [])
assert url == "https://api.minimax.io/v2/query/video_generation/a%3Fview%3Dadmin%2F..%2Fx"
row.descriptor = {"poll": {"endpoint": "minimax.video-gen.task.status",
"param": {"in": "pathParams", "name": "task_id"}}}
with pytest.raises(RuntimeError):
task_app._poll_target(row) # declared as a path parameter, but the target has no placeholder
def test_fetch_command_and_shown_neutralise_provider_strings():
assert asynctasks.fetch_command({"endpoint": "minimax.video-gen.result.retrieve", "name": "file_id",
"value": "x; touch /tmp/pwned"}) == \
"treg call minimax.video-gen.result.retrieve -p 'file_id=x; touch /tmp/pwned'"
assert asynctasks.shown("ok-123") == "ok-123"
assert asynctasks.shown("id\nresume: treg call evil") == "id\\nresume: treg call evil"
assert asynctasks.shown("\x1b]52;c;aGk=\x07") == "\\x1b]52;c;aGk=\\x07"
async def test_idempotent_replay_of_an_async_submission_keeps_the_descriptor(
clients: AsyncClient, monkeypatch, replicate_platform,
):
async def fake_relay(*args, **kwargs):
return _response(201, {"id": "prediction-idem", "urls": {"get": "https://api.replicate.com/v1/predictions/i"}})
monkeypatch.setattr(call_service, "relay", fake_relay)
body = {"input": {"prompt": "A kite.", "num_outputs": 1, "aspect_ratio": "1:1", "output_format": "webp"}}
first = await clients.post(f"/call/{EP}", json=body, headers={"Idempotency-Key": "gen-1"})
assert first.status_code == 201 and "x-treg-async" in {k.lower() for k in first.headers}
again = await clients.post(f"/call/{EP}", json=body, headers={"Idempotency-Key": "gen-1"})
assert again.status_code == 201 and again.headers.get("X-Treg-Idempotent-Replay") == "true"
assert again.headers.get("x-treg-async") == first.headers.get("x-treg-async")
async def test_two_workers_racing_the_same_row_move_money_exactly_once(
clients: AsyncClient, monkeypatch, replicate_platform,
):
"""A second instance re-claims a row whose lease lapsed while the first poll is in flight
(Render cron overlap). Both reach _finish; the row lock and the once-only hold claim leave one
settle entry and one terminal state. Meaningful on Postgres (FOR UPDATE SKIP LOCKED)."""
call_id = await _due_submission(clients, monkeypatch, {"status": "succeeded", "output": ["u"]})
first_polling = asyncio.Event()
release_first = asyncio.Event()
polls = 0
async def slow_poll(row, client):
nonlocal polls
polls += 1
if polls == 1:
async with session_maker() as db: # the lease lapses while this poll is in flight
live = await db.get(AsyncTaskRecord, row.call_id)
live.next_check_at = utcnow_naive() - timedelta(seconds=1)
await db.commit()
first_polling.set()
await release_first.wait()
return 200, json.dumps({"status": "succeeded", "output": ["u"]}).encode()
monkeypatch.setattr(task_app, "_poll", slow_poll)
first = asyncio.create_task(task_app.settle_due())
await asyncio.wait_for(first_polling.wait(), 10)
second = await task_app.settle_due() # re-claims the lapsed lease and settles
release_first.set()
first_result = await first
assert second.claimed == 1 and first_result.claimed == 1
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
settles = (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == call_id, LedgerEntry.kind == "settle"))).scalars().all()
hold = await db.get(Hold, call_id)
assert row.status == "settled" and row.settled_micro == 3000
assert len(settles) == 1 and settles[0].amount_micro == -3000 and hold is None
async def test_cancellation_at_the_pending_row_commit_boundary_leaves_a_coherent_outcome(
clients: AsyncClient, monkeypatch, replicate_platform,
):
"""The request is cancelled the instant the pending row commits: the request path releases the
hold it still owns, and the worker must then record the row as released, not settle it at zero."""
real_defer = task_app.defer_submission
async def defer_then_cancel(mk, body, org_id, *, tags=None):
await real_defer(mk, body, org_id, tags=tags)
mk.call_id = mk.call_id or None
raise asyncio.CancelledError()
monkeypatch.setattr(task_app, "defer_submission", defer_then_cancel)
# The request path re-raises CancelledError after compensating; the ASGI client surfaces it.
with pytest.raises(BaseException):
await _submit(clients, monkeypatch, {
"id": "prediction-cancel", "urls": {"get": "https://api.replicate.com/v1/predictions/c"}})
async with session_maker() as db:
row = (await db.execute(select(AsyncTaskRecord).where(
AsyncTaskRecord.task_id == "prediction-cancel"))).scalar_one()
hold = await db.get(Hold, row.call_id)
kinds = sorted(e.kind for e in (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == row.call_id))).scalars().all())
row.next_check_at = utcnow_naive() - timedelta(seconds=1)
await db.commit()
call_id = row.call_id
assert row.status == "pending"
if hold is None: # compensation released the hold while the row was already durable
assert kinds == ["release", "reserve"]
async def fake_poll(row, client):
return 200, json.dumps({"status": "succeeded", "output": ["u"]}).encode()
monkeypatch.setattr(task_app, "_poll", fake_poll)
await task_app.settle_due()
async with session_maker() as db:
row = await db.get(AsyncTaskRecord, call_id)
kinds = sorted(e.kind for e in (await db.execute(select(LedgerEntry).where(
LedgerEntry.call_id == call_id))).scalars().all())
assert row.status == "released" and row.settled_micro == 0
assert kinds == ["release", "reserve"], "no settle may follow a released hold"
else: # the hold survived with its row: the worker owns it from here, nothing was double-moved
assert kinds == ["reserve"]
async def test_another_org_cannot_read_a_task_by_its_call_ref(
clients: AsyncClient, monkeypatch, replicate_platform,
):
call_id = await _due_submission(clients, monkeypatch, {"status": "succeeded", "output": ["u"]})
assert (await task_app.settle_due()).settled == 1
assert call_id in await task_app.views_for(1, [call_id])
assert await task_app.views_for(2, [call_id]) == {}
other = await clients.post("/users", json={"email": "someone-else@example.com"})
assert other.status_code == 200
stranger = {"X-Treg-Token": other.json()["token"]}
assert (await clients.get(f"/calls/{call_id}", headers=stranger)).status_code == 404
assert call_id not in {r.get("call_ref") for r in (await clients.get("/calls", headers=stranger)).json()}
def _upstream_idempotency_keys(relayed: list) -> list[str]:
return [v.decode() for req in relayed for k, v in req.raw_headers if k.lower() == b"idempotency-key"]
async def test_shared_key_idempotency_label_is_partitioned_per_org(
clients: AsyncClient, monkeypatch, replicate_platform,
):
"""Two orgs sending one Idempotency-Key on treg's key must not collide on the provider account.
Reproduced live against LeadsForge (2026-09-09): the provider returned org A's job to org B under
the shared label, and `resource_ownership.produces` then made B its owner.
"""
relayed = []
async def fake_relay(request, *args, **kwargs):
relayed.append(request)
return _response(201, {"id": f"prediction-{len(relayed)}", "status": "starting"})
monkeypatch.setattr(call_service, "relay", fake_relay)
body = {"input": {"prompt": "A red kite over a beach.", "num_outputs": 1,
"aspect_ratio": "1:1", "output_format": "webp"}}
label = {"Idempotency-Key": "retry-1"}
other = await verified_signup(clients, json={"email": "idem-stranger@example.com"})
stranger = {"X-Treg-Token": other.json()["token"], **label}
assert (await clients.post(f"/call/{EP}", json=body, headers=label)).status_code == 201
assert (await clients.post(f"/call/{EP}", json=body, headers=stranger)).status_code == 201
keys = _upstream_idempotency_keys(relayed)
assert len(keys) == 2 and keys[0] != keys[1], "the two orgs reached the provider under one label"
assert "retry-1" not in keys, "the caller's raw label reached the shared provider account"
# The same org retrying the same label is still served by treg's own replay, not the provider.
assert (await clients.post(f"/call/{EP}", json=body, headers=label)).status_code == 201
assert len(relayed) == 2
async def test_own_key_relays_idempotency_label_verbatim(clients: AsyncClient, monkeypatch):
await clients.post("/secrets", json={"name": "replicate", "value": "own-token"})
relayed = []
async def fake_relay(request, *args, **kwargs):
relayed.append(request)
return _response(201, {"id": "own-account-prediction", "status": "starting"})
monkeypatch.setattr(call_service, "relay", fake_relay)
response = await clients.post(f"/call/{EP}", json={"input": {"prompt": "x"}},
headers={"Idempotency-Key": "retry-1"})
assert response.status_code == 201
assert _upstream_idempotency_keys(relayed) == ["retry-1"]