Address EHBP review comments

This commit is contained in:
9qeklajc
2026-07-01 22:41:14 +02:00
parent 07d39c2a7b
commit 7328a5ac45
3 changed files with 196 additions and 10 deletions
+11 -7
View File
@@ -199,6 +199,9 @@ def _resolve_ehbp_target_url(
def _validated_confidential_target_url(
enclave_url: str, profile: "ConfidentialInferenceProfile"
) -> str | None:
# Client target overrides are Tinfoil-only for now. Future confidential
# inference providers must add their own constrained validator here before
# opting into ``allow_client_target_override``.
if profile.client_target_url_header == _ENCLAVE_URL_HEADER:
return _validated_tinfoil_enclave_base_url(enclave_url)
return None
@@ -370,12 +373,9 @@ def _extract_usage_from_response(
class ConfidentialInferenceProfile:
"""Provider-neutral policy for encrypted/confidential inference forwarding."""
protocol: str = "EHBP"
usage_response_header: str | None = None
client_target_url_header: str | None = None
allow_client_target_override: bool = False
trusted_model_binding_header: str | None = None
missing_usage_billing_policy: str = "max_cost"
proxy_only_headers: frozenset[str] = _PROXY_ONLY_HEADERS
@@ -397,6 +397,8 @@ async def finalize_ehbp_actual_cost_payment(
) -> None:
"""Finalize an EHBP bearer request using clamped provider usage metrics."""
billing_key = await get_billing_key(key, session)
key_hash = key.hashed_key
billing_key_hash = billing_key.hashed_key
total_cost_msats = max(0, int(cost_info.get("total_msats", reserved_cost_for_model)))
now = int(time.time())
@@ -445,8 +447,8 @@ async def finalize_ehbp_actual_cost_payment(
logger.error(
"Failed to finalize EHBP usage-based payment",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"key_hash": key_hash[:8] + "...",
"billing_key_hash": billing_key_hash[:8] + "...",
"model": model_id,
"reserved_cost_for_model": reserved_cost_for_model,
"total_cost_msats": total_cost_msats,
@@ -504,6 +506,8 @@ async def finalize_ehbp_max_cost_payment(
cost and releases the reservation.
"""
billing_key = await get_billing_key(key, session)
key_hash = key.hashed_key
billing_key_hash = billing_key.hashed_key
total_cost_msats = max(0, int(max_cost_for_model))
now = int(time.time())
@@ -567,8 +571,8 @@ async def finalize_ehbp_max_cost_payment(
logger.error(
"Failed to finalize EHBP max-cost payment",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"key_hash": key_hash[:8] + "...",
"billing_key_hash": billing_key_hash[:8] + "...",
"model": model_id,
"max_cost_for_model": max_cost_for_model,
"parent_rowcount": result.rowcount,
-3
View File
@@ -60,12 +60,9 @@ class TinfoilUpstreamProvider(BaseUpstreamProvider):
platform_url = "https://docs.tinfoil.sh"
supports_ehbp = True
confidential_inference_profile = ConfidentialInferenceProfile(
protocol="EHBP",
usage_response_header=_RESPONSE_USAGE_HEADER,
client_target_url_header=_ENCLAVE_URL_HEADER,
allow_client_target_override=True,
trusted_model_binding_header=None,
missing_usage_billing_policy="max_cost",
proxy_only_headers=_PROXY_ONLY_HEADERS,
)
+185
View File
@@ -0,0 +1,185 @@
from __future__ import annotations
from typing import AsyncGenerator
import pytest
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlalchemy.pool import StaticPool
from sqlmodel import SQLModel, select
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.db import ApiKey
from routstr.upstream.ehbp import (
finalize_ehbp_actual_cost_payment,
finalize_ehbp_max_cost_payment,
)
def _make_engine() -> AsyncEngine:
return create_async_engine(
"sqlite+aiosqlite://",
poolclass=StaticPool,
connect_args={"check_same_thread": False},
)
@pytest.fixture
async def session(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[AsyncSession, None]:
monkeypatch.setattr("routstr.upstream.ehbp.ROUTSTR_FEE_PERCENT", 0)
engine = _make_engine()
async with engine.begin() as conn:
await conn.run_sync(SQLModel.metadata.create_all)
db_session = AsyncSession(engine, expire_on_commit=False)
try:
yield db_session
finally:
await db_session.close()
await engine.dispose()
async def _api_key(session: AsyncSession, hashed_key: str) -> ApiKey | None:
return (
await session.exec(select(ApiKey).where(ApiKey.hashed_key == hashed_key))
).one_or_none()
@pytest.mark.asyncio
async def test_finalize_actual_cost_payment_updates_balance_and_releases_reserve(
session: AsyncSession,
) -> None:
key = ApiKey(
hashed_key="ehbp-actual",
balance=10_000,
reserved_balance=3_000,
reserved_at=123,
)
session.add(key)
await session.commit()
await finalize_ehbp_actual_cost_payment(
key,
session,
reserved_cost_for_model=3_000,
model_id="tinfoil/model",
cost_info={
"total_msats": 1_200,
"input_tokens": 10,
"output_tokens": 20,
"input_msats": 500,
"output_msats": 700,
},
)
updated = await _api_key(session, "ehbp-actual")
assert updated is not None
assert updated.balance == 8_800
assert updated.reserved_balance == 0
assert updated.reserved_at is None
assert updated.total_spent == 1_200
@pytest.mark.asyncio
async def test_finalize_max_cost_payment_updates_parent_and_child_spend(
session: AsyncSession,
) -> None:
parent = ApiKey(
hashed_key="ehbp-parent",
balance=10_000,
reserved_balance=3_000,
reserved_at=123,
)
child = ApiKey(
hashed_key="ehbp-child",
balance=0,
reserved_balance=3_000,
reserved_at=123,
parent_key_hash="ehbp-parent",
)
session.add(parent)
session.add(child)
await session.commit()
await finalize_ehbp_max_cost_payment(
child,
session,
max_cost_for_model=3_000,
model_id="tinfoil/model",
)
updated_parent = await _api_key(session, "ehbp-parent")
updated_child = await _api_key(session, "ehbp-child")
assert updated_parent is not None
assert updated_child is not None
assert updated_parent.balance == 7_000
assert updated_parent.reserved_balance == 0
assert updated_parent.reserved_at is None
assert updated_parent.total_spent == 3_000
assert updated_child.balance == 0
assert updated_child.reserved_balance == 0
assert updated_child.reserved_at is None
assert updated_child.total_spent == 3_000
@pytest.mark.asyncio
async def test_finalize_actual_cost_payment_rolls_back_when_parent_update_matches_no_rows(
session: AsyncSession,
) -> None:
key = ApiKey(
hashed_key="ehbp-missing-parent",
balance=10_000,
reserved_balance=3_000,
reserved_at=123,
)
session.add(key)
await session.commit()
await session.delete(key)
await session.commit()
await finalize_ehbp_actual_cost_payment(
key,
session,
reserved_cost_for_model=3_000,
model_id="tinfoil/model",
cost_info={"total_msats": 1_200},
)
assert await _api_key(session, "ehbp-missing-parent") is None
@pytest.mark.asyncio
async def test_finalize_max_cost_payment_rolls_back_parent_when_child_update_matches_no_rows(
session: AsyncSession,
) -> None:
parent = ApiKey(
hashed_key="ehbp-rollback-parent",
balance=10_000,
reserved_balance=3_000,
reserved_at=123,
)
child = ApiKey(
hashed_key="ehbp-missing-child",
balance=0,
reserved_balance=3_000,
reserved_at=123,
parent_key_hash="ehbp-rollback-parent",
)
session.add(parent)
session.add(child)
await session.commit()
await session.delete(child)
await session.commit()
await finalize_ehbp_max_cost_payment(
child,
session,
max_cost_for_model=3_000,
model_id="tinfoil/model",
)
updated_parent = await _api_key(session, "ehbp-rollback-parent")
assert updated_parent is not None
assert updated_parent.balance == 10_000
assert updated_parent.reserved_balance == 3_000
assert updated_parent.reserved_at == 123
assert updated_parent.total_spent == 0
assert await _api_key(session, "ehbp-missing-child") is None