mirror of
https://github.com/superdesigndev/treg.git
synced 2026-10-02 03:24:35 +08:00
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.
1507 lines
72 KiB
Python
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"]
|