Files

498 lines
19 KiB
Python

"""Tests for provider-neutral account usage and the paired-device endpoint."""
from __future__ import annotations
from datetime import datetime, timezone
from pathlib import Path
from types import SimpleNamespace
import tempfile
import unittest
from unittest import mock
from aiohttp import web
from aiohttp.test_utils import AioHTTPTestCase
from plugin.relay.config import RelayConfig
from plugin.relay.provider_usage import (
collect_provider_usage,
fetch_codex_usage,
fetch_nous_usage,
fetch_opencode_go_usage,
fetch_supergrok_usage,
resolve_profile_home,
serialize_account_snapshot,
unavailable_provider,
)
from plugin.relay.active_credentials import record_active_credential
from plugin.relay.server import create_app
class _FakeResponse:
def __init__(self, status: int = 200, payload: dict | None = None):
self.status = status
self._payload = payload or {}
async def __aenter__(self):
return self
async def __aexit__(self, *exc):
return False
async def json(self):
return self._payload
class _FakeSession:
def __init__(self, response: _FakeResponse):
self.response = response
self.headers: dict | None = None
async def __aenter__(self):
return self
async def __aexit__(self, *exc):
return False
def get(self, _url, *, headers=None, timeout=None):
self.headers = headers
return self.response
class _SequencedSession:
"""Yield queued responses in call order, like aiohttp's request context manager."""
def __init__(self, responses: list[_FakeResponse]):
self._responses = list(responses)
self.calls: list[dict] = []
async def __aenter__(self):
return self
async def __aexit__(self, *exc):
return False
def get(self, url, *, headers=None, timeout=None):
self.calls.append({"url": url, "headers": headers or {}})
if not self._responses:
raise AssertionError("unexpected extra provider request")
return self._responses.pop(0)
class ProviderUsageModelTests(unittest.IsolatedAsyncioTestCase):
def test_profile_home_is_exact_and_rejects_traversal(self) -> None:
with tempfile.TemporaryDirectory() as raw:
# Resolve the temp root so macOS /var -> /private/var matches
# Path.resolve() inside resolve_profile_home.
root = Path(raw).resolve()
(root / "config.yaml").write_text("model: {}\n", encoding="utf-8")
victor = (root / "profiles" / "victor").resolve()
victor.mkdir(parents=True)
(victor / "config.yaml").write_text("model: {}\n", encoding="utf-8")
self.assertEqual(resolve_profile_home(str(root / "config.yaml"), "Victor"), victor)
with self.assertRaises(ValueError):
resolve_profile_home(str(root / "config.yaml"), "../victor")
def test_serializes_upstream_snapshot_without_credentials(self) -> None:
snapshot = SimpleNamespace(
available=True,
source="usage_api",
fetched_at=datetime(2026, 8, 21, tzinfo=timezone.utc),
plan="Plus",
windows=(
SimpleNamespace(
label="Session",
used_percent=42.5,
reset_at=datetime(2026, 8, 22, tzinfo=timezone.utc),
detail=None,
),
),
details=("Credits balance: $4.20",),
)
result = serialize_account_snapshot(
snapshot,
provider_id="openai-codex",
display_name="Codex",
)
self.assertEqual(result["status"], "available")
self.assertEqual(result["windows"][0]["used_percent"], 42.5)
self.assertEqual(result["plan"], "Plus")
self.assertNotIn("token", result)
async def test_opencode_missing_key_is_not_configured(self) -> None:
result = await fetch_opencode_go_usage(credential_resolver=lambda _provider: {})
self.assertEqual(result["id"], "opencode-go")
self.assertEqual(result["status"], "not_configured")
async def test_nous_exposes_structured_balances_without_raw_mobile_details(self) -> None:
account = SimpleNamespace(
logged_in=True,
paid_service_access=True,
paid_service_access_info=SimpleNamespace(
subscription_credits_remaining=31.98,
purchased_credits_remaining=0.0,
total_usable_credits=31.98,
),
subscription=SimpleNamespace(
plan="Plus",
monthly_credits=None,
credits_remaining=31.98,
rollover_credits=10.0,
current_period_end="2026-09-18T00:11:42.000Z",
),
portal_base_url="https://portal.nousresearch.com",
org_slug="example",
)
result = await fetch_nous_usage(
account_fetcher=lambda **_kwargs: account,
)
self.assertEqual(result["status"], "available")
self.assertEqual(result["plan"], "Plus")
self.assertEqual(result["balances"][0], {
"id": "total",
"label": "Total usable",
"amount": 31.98,
"currency": "USD",
})
self.assertEqual(result["renews_at"], "2026-09-18T00:11:42.000Z")
self.assertTrue(result["action_url"].endswith("/orgs/example/billing?topup=open"))
self.assertEqual(result["details"], [])
async def test_opencode_normalizes_windows_without_inventing_dollars(self) -> None:
fake = _FakeSession(
_FakeResponse(
payload={
"usage": {
"rolling": {"percent": 42, "resetsAt": "2026-08-22T00:00:00Z"},
"weekly": {"percent": 18},
}
}
)
)
result = await fetch_opencode_go_usage(
session_factory=lambda: fake,
credential_resolver=lambda _provider: {
"api_key": "secret",
"base_url": "https://opencode.ai/zen/go/v1",
},
)
self.assertEqual(result["status"], "available")
self.assertEqual([row["id"] for row in result["windows"]], ["rolling", "weekly"])
self.assertNotIn("limits", result)
self.assertEqual(fake.headers["Authorization"], "Bearer secret")
async def test_supergrok_without_oauth_is_not_configured(self) -> None:
result = await fetch_supergrok_usage(credential_resolver=lambda: {})
self.assertEqual(result["id"], "supergrok")
self.assertEqual(result["status"], "not_configured")
self.assertEqual(result["windows"], [])
async def test_supergrok_missing_oauth_state_is_not_configured(self) -> None:
class MissingOAuthState(Exception):
code = "xai_auth_missing"
def resolve_credentials() -> dict:
raise MissingOAuthState("No credentials stored")
result = await fetch_supergrok_usage(credential_resolver=resolve_credentials)
self.assertEqual(result["status"], "not_configured")
async def test_supergrok_oauth_refresh_failure_is_unavailable(self) -> None:
class RefreshFailure(Exception):
code = "xai_refresh_failed"
def resolve_credentials() -> dict:
raise RefreshFailure("private token details")
result = await fetch_supergrok_usage(credential_resolver=resolve_credentials)
self.assertEqual(result["status"], "unavailable")
self.assertEqual(result["message"], "Could not resolve SuperGrok credentials")
self.assertNotIn("private token details", str(result))
async def test_supergrok_maps_subscription_and_product_windows(self) -> None:
session = _SequencedSession(
[
_FakeResponse(payload={"userId": "user-1"}),
_FakeResponse(
payload={
"subscriptionTier": "SuperGrok",
"onDemandEnabled": True,
"config": {
"creditUsagePercent": 14,
"currentPeriod": {
"type": "USAGE_PERIOD_TYPE_WEEKLY",
"start": "2026-09-06T08:34:12.348291+00:00",
"end": "2026-09-13T08:34:12.348291+00:00",
},
"productUsage": [
{"product": "GrokBuild", "usagePercent": 11},
{"product": "GrokImagine", "usagePercent": 2},
{"product": "GrokChat", "usagePercent": None},
],
"onDemandCap": {"val": 500},
"onDemandUsed": {"val": 125},
"prepaidBalance": {"val": 0},
},
}
),
]
)
result = await fetch_supergrok_usage(
session_factory=lambda: session,
credential_resolver=lambda: {"api_key": "secret"},
)
self.assertEqual(result["status"], "available")
self.assertEqual(result["source"], "provider_api")
self.assertEqual(result["plan"], "SuperGrok")
self.assertEqual(
[row["id"] for row in result["windows"]],
["period", "product_grok_build", "product_grok_imagine"],
)
self.assertEqual(result["windows"][0]["label"], "Weekly")
self.assertEqual(result["windows"][0]["used_percent"], 14.0)
self.assertEqual(result["windows"][0]["reset_at"], "2026-09-13T08:34:12.348291+00:00")
self.assertEqual(result["windows"][1]["label"], "Grok Build")
self.assertEqual(result["windows"][1]["used_percent"], 11.0)
self.assertEqual(result["details"], ["On-demand: $1.25 used of $5.00"])
self.assertEqual(session.calls[0]["headers"]["Authorization"], "Bearer secret")
self.assertEqual(session.calls[1]["headers"]["x-userid"], "user-1")
self.assertIn("/billing?format=credits", session.calls[1]["url"])
self.assertNotIn("secret", str(result))
async def test_supergrok_stops_before_billing_without_account_identity(self) -> None:
session = _SequencedSession([_FakeResponse(payload={"userId": ""})])
result = await fetch_supergrok_usage(
session_factory=lambda: session,
credential_resolver=lambda: {"api_key": "secret"},
)
self.assertEqual(result["status"], "unavailable")
self.assertEqual(len(session.calls), 1)
self.assertNotIn("secret", str(result))
async def test_supergrok_reports_top_level_on_demand_state_without_amounts(self) -> None:
session = _SequencedSession(
[
_FakeResponse(payload={"userId": "user-1"}),
_FakeResponse(
payload={
"onDemandEnabled": True,
"config": {
"creditUsagePercent": 0,
"currentPeriod": {"type": "USAGE_PERIOD_TYPE_WEEKLY"},
},
}
),
]
)
result = await fetch_supergrok_usage(
session_factory=lambda: session,
credential_resolver=lambda: {"api_key": "secret"},
)
self.assertEqual(result["status"], "available")
self.assertEqual(result["details"], ["On-demand enabled"])
async def test_supergrok_fresh_period_surfaces_window_without_inventing_a_percent(self) -> None:
session = _SequencedSession(
[
_FakeResponse(payload={"userId": "user-1"}),
_FakeResponse(
payload={
"config": {
"currentPeriod": {
"type": "USAGE_PERIOD_TYPE_WEEKLY",
"start": "2026-09-13T08:34:12.348291+00:00",
"end": "2026-09-20T08:34:12.348291+00:00",
},
"billingPeriodEnd": "2026-09-20T08:34:12.348291+00:00",
}
}
),
]
)
result = await fetch_supergrok_usage(
session_factory=lambda: session,
credential_resolver=lambda: {"api_key": "secret"},
)
self.assertEqual(result["status"], "available")
self.assertEqual(len(result["windows"]), 1)
self.assertEqual(result["windows"][0]["label"], "Weekly")
self.assertIsNone(result["windows"][0]["used_percent"])
self.assertEqual(result["windows"][0]["reset_at"], "2026-09-20T08:34:12.348291+00:00")
self.assertEqual(result["windows"][0]["detail"], "No usage reported yet")
async def test_supergrok_unusable_payload_is_unavailable(self) -> None:
session = _SequencedSession(
[
_FakeResponse(payload={"userId": "user-1"}),
_FakeResponse(payload={"config": {"isUnifiedBillingUser": True}}),
]
)
result = await fetch_supergrok_usage(
session_factory=lambda: session,
credential_resolver=lambda: {"api_key": "secret"},
)
self.assertEqual(result["status"], "unavailable")
self.assertEqual(result["windows"], [])
self.assertEqual(result["message"], "Provider returned no usage windows")
async def test_collection_keeps_provider_order_and_schema(self) -> None:
async def codex(_home, **_kwargs):
return unavailable_provider("openai-codex", "Codex")
async def nous(_home):
return unavailable_provider("nous", "Nous")
async def opencode(*, profile_home=None):
return unavailable_provider("opencode-go", "OpenCode Go")
async def supergrok(*, profile_home=None):
return unavailable_provider("supergrok", "SuperGrok")
result = await collect_provider_usage(
codex_fetcher=codex,
nous_fetcher=nous,
opencode_fetcher=opencode,
supergrok_fetcher=supergrok,
)
self.assertEqual(result["schema_version"], 2)
self.assertEqual(
result["capabilities"],
["credential_pools", "structured_balances", "opencode_go", "supergrok"],
)
self.assertEqual(
[row["id"] for row in result["providers"]],
["openai-codex", "nous", "opencode-go", "supergrok"],
)
async def test_codex_pool_marks_exact_live_session_credential_active(self) -> None:
with tempfile.TemporaryDirectory() as raw:
home = Path(raw)
record_active_credential(
home,
session_id="session-2",
provider_id="openai-codex",
credential_id="entry-2",
)
entries = [
SimpleNamespace(
id=f"entry-{index}",
label=f"Account {index}",
last_status="ok",
last_status_at=None,
last_error_reset_at=None,
runtime_base_url="https://chatgpt.com/backend-api/codex",
runtime_api_key=f"secret-{index}",
)
for index in (1, 2)
]
snapshots = {
"secret-1": SimpleNamespace(
available=True,
source="usage_api",
fetched_at=datetime(2026, 8, 21, tzinfo=timezone.utc),
plan="Pro",
windows=(SimpleNamespace(label="Session", used_percent=100, reset_at=None, detail=None),),
details=(),
),
"secret-2": SimpleNamespace(
available=True,
source="usage_api",
fetched_at=datetime(2026, 8, 21, tzinfo=timezone.utc),
plan="Pro",
windows=(SimpleNamespace(label="Session", used_percent=24, reset_at=None, detail=None),),
details=(),
),
}
result = await fetch_codex_usage(
home,
session_id="session-2",
pool_loader=lambda _provider: SimpleNamespace(entries=lambda: entries),
snapshot_fetcher=lambda *, api_key, base_url: snapshots[api_key],
)
self.assertEqual(result["active_credential_state"], "known")
active = next(row for row in result["credentials"] if row["active"])
self.assertEqual(active["label"], "Account 2")
self.assertEqual(active["windows"][0]["used_percent"], 24.0)
limited = next(row for row in result["credentials"] if row["label"] == "Account 1")
self.assertEqual(limited["status"], "at_limit")
self.assertNotIn("secret-2", str(result))
class ProviderUsageEndpointTests(AioHTTPTestCase):
async def get_application(self) -> web.Application:
return create_app(RelayConfig(provider_usage_enabled=self.usage_enabled))
@property
def usage_enabled(self) -> bool:
return True
def _server(self):
return self.app["server"]
async def _mint(self) -> str:
return self._server().sessions.create_session("phone", "device").token
async def test_requires_bearer(self) -> None:
response = await self.client.get("/usage/providers")
self.assertEqual(response.status, 401)
async def test_invalid_profile_error_does_not_reflect_request_input(self) -> None:
token = await self._mint()
response = await self.client.get(
"/usage/providers?profile=../private-token",
headers={"Authorization": f"Bearer {token}"},
)
body = await response.text()
self.assertEqual(response.status, 400 if self.usage_enabled else 404)
if self.usage_enabled:
self.assertEqual(body, "invalid or unknown profile")
self.assertNotIn("private-token", body)
@mock.patch(
"plugin.relay.server.collect_provider_usage",
new=mock.AsyncMock(return_value={"schema_version": 1, "providers": []}),
)
async def test_returns_normalized_payload(self) -> None:
token = await self._mint()
response = await self.client.get(
"/usage/providers",
headers={"Authorization": f"Bearer {token}"},
)
self.assertEqual(response.status, 200)
self.assertEqual((await response.json())["schema_version"], 1)
class ProviderUsageDisabledEndpointTests(ProviderUsageEndpointTests):
@property
def usage_enabled(self) -> bool:
return False
async def test_returns_normalized_payload(self) -> None:
token = await self._mint()
response = await self.client.get(
"/usage/providers",
headers={"Authorization": f"Bearer {token}"},
)
self.assertEqual(response.status, 404)
if __name__ == "__main__":
unittest.main()