feat: add cached pricing workspace

This commit is contained in:
2026-08-04 00:16:34 +08:00
parent 81e9ce1530
commit 311236d28f
5 changed files with 1156 additions and 376 deletions
+420 -10
View File
@@ -1,13 +1,17 @@
from __future__ import annotations
import contextlib
import copy
import io
import os
from pathlib import Path
import importlib.util
import sys
import tempfile
import threading
import time
import unittest
from unittest import mock
def load_module():
@@ -20,6 +24,117 @@ def load_module():
return module
def workspace_payload_fixture() -> dict[str, object]:
return {
"service": "sub2api-pricing-monitor",
"view": "workspace",
"generated_at": "2026-08-03T12:00:00Z",
"state": {"ok": True, "partial": False},
"components": {
"accounts": {"ok": True, "stale": False},
"status": {"ok": True, "stale": False},
"traffic": {"ok": True, "stale": False},
},
"sources": [
{
"name": "code-plan",
"source_kind": "newapi",
"health_state": "healthy",
"balance_available": True,
"balance": {"available": 5000000, "available_cny": 10, "unit": "quota"},
}
],
"accounts": {
"source_name": "workspace",
"totals": {
"total_accounts": 1,
"usable_accounts": 1,
"today_cost_usd": 1.25,
"today_tokens": 300,
"today_requests": 2,
},
"accounts": [
{
"id": "9007199254740993",
"name": "oai-quota-code-plan",
"routing_group": "fast",
"provider": "openai",
"kind": "quota_limited",
"usable": True,
}
],
},
"status": {
"channel_monitors": {
"items": [
{
"id": "9007199254740995",
"name": "oai-quota-code-plan",
"provider": "openai",
"latest_status": "success",
}
]
}
},
"traffic": {
"window_hours": 24,
"sample_limit": 1000,
"sample_limited": False,
"partial": False,
"instances": [
{"instance": "server6", "ok": True, "error_total": 4},
{"instance": "server4", "ok": True, "error_total": 0},
],
"requests": [
{
"instance": "server6",
"id": "9007199254740997",
"created_at": "2026-08-03T11:59:00Z",
"api_key_id": "7",
"api_key_name": "wmy",
"account_id": "9",
"account_name": "oai-quota-code-plan",
"model": "gpt-5.5",
"input_tokens": 100,
"output_tokens": 50,
"cache_read_tokens": 25,
"actual_cost": 0.25,
"duration_ms": 1000,
}
],
"errors": [
{
"instance": "server6",
"latest_at": "2026-08-03T11:58:00Z",
"error_count": 4,
"status_code": 502,
"inbound_status_code": 500,
"upstream_status_code": 502,
"api_key_id": "7",
"api_key_name": "wmy",
"account_id": "9",
"account_name": "oai-quota-code-plan",
"model": "gpt-5.5",
"error_type": "api_error",
"error_source": "upstream_http",
}
],
"keys": [
{
"instance": "server6",
"api_key_id": "7",
"api_key_name": "wmy",
"status": "active",
"request_count": 2,
"token_count": 175,
"actual_cost": 0.25,
"latest_at": "2026-08-03T11:59:00Z",
}
],
},
}
class Sub2APIQuotaTUITests(unittest.TestCase):
def test_normalize_rows_sorts_by_usage_and_formats_windows(self) -> None:
mod = load_module()
@@ -190,7 +305,7 @@ class Sub2APIQuotaTUITests(unittest.TestCase):
self.assertIsNone(unavailable_rows[0]["balance"])
self.assertIsNone(unavailable_rows[0]["balance_cny"])
self.assertEqual([row["name"] for row in mod.normalize_pricing_rows(payload, "keday")], ["kedaya"])
self.assertEqual(mod.pricing_summary(payload, rows), "upstreams 1/2 healthy | CNY ¥10")
self.assertEqual(mod.pricing_summary(payload, rows), "sources 1/2 healthy | CNY ¥10")
out = io.StringIO()
with contextlib.redirect_stdout(out):
mod.print_pricing_once(payload)
@@ -314,8 +429,8 @@ class Sub2APIQuotaTUITests(unittest.TestCase):
def __exit__(self, exc_type, exc, traceback):
return False
def read(self):
return compressed
def read(self, size=-1):
return compressed if size < 0 else compressed[:size]
old_urlopen = mod.urllib.request.urlopen
mod.urllib.request.urlopen = lambda request, timeout: (requests.append(request) or FakeResponse())
@@ -330,13 +445,35 @@ class Sub2APIQuotaTUITests(unittest.TestCase):
class PlainResponse:
headers = {}
def read(self):
return mod.json.dumps(payload).encode("utf-8")
def read(self, size=-1):
raw = mod.json.dumps(payload).encode("utf-8")
return raw if size < 0 else raw[:size]
self.assertEqual(mod.decode_json_response(PlainResponse(), "plain JSON failed"), payload)
self.assertEqual(len(requests), 4)
self.assertTrue(all(request.get_header("Accept-encoding") == "gzip" for request in requests))
def test_json_decoder_bounds_plain_and_gzip_payloads(self) -> None:
mod = load_module()
class Response:
def __init__(self, raw: bytes, encoding: str = "") -> None:
self.raw = raw
self.headers = {"Content-Encoding": encoding} if encoding else {}
def read(self, size=-1):
return self.raw if size < 0 else self.raw[:size]
oversized = mod.json.dumps({"value": "x" * 256}).encode("utf-8")
with self.assertRaisesRegex(RuntimeError, "bounded"):
mod.decode_json_response(Response(oversized), "bounded response", maximum_bytes=64)
with self.assertRaisesRegex(RuntimeError, "bounded"):
mod.decode_json_response(
Response(mod.gzip.compress(oversized), "gzip"),
"bounded response",
maximum_bytes=64,
)
def test_default_refresh_intervals_are_five_minutes(self) -> None:
mod = load_module()
@@ -613,6 +750,181 @@ def sample_logs_payload() -> dict:
}
class WorkspaceTests(unittest.TestCase):
def test_workspace_adapters_preserve_ids_and_map_compact_traffic(self) -> None:
mod = load_module()
payload = workspace_payload_fixture()
accounts = mod.workspace_accounts_payload(payload)
status = mod.workspace_status_payload(payload)
pricing = mod.workspace_pricing_payload(payload)
logs = mod.workspace_logs_payload(payload)
keys = mod.workspace_keys_payload(payload)
errors = mod.workspace_errors_payload(payload)
self.assertEqual(accounts["accounts"][0]["id"], "9007199254740993")
self.assertEqual(status["channel_monitors"]["items"][0]["id"], "9007199254740995")
self.assertEqual(pricing["sources"][0]["name"], "code-plan")
self.assertEqual(logs["data"]["items"][0]["id"], "9007199254740997")
self.assertEqual(logs["data"]["items"][0]["instance"], "server6")
normalized_accounts = mod.normalize_account_rows(accounts, pricing_payload=pricing)
normalized_logs = mod.normalize_log_rows(logs)
self.assertEqual(normalized_accounts[0]["id"], 9007199254740993)
self.assertEqual(normalized_logs[0]["id"], 9007199254740997)
self.assertEqual(normalized_logs[0]["cost"], 0.25)
self.assertEqual(mod.as_int("9007199254740999"), 9007199254740999)
self.assertEqual(mod.normalize_key_rows(keys)[0]["name"], "wmy")
self.assertEqual(mod.normalize_key_rows(keys)[0]["cost"], 0.25)
self.assertEqual(errors["items"][0]["error_count"], 4)
self.assertEqual(errors["items"][0]["phase"], "upstream_http")
self.assertEqual(errors["sources"]["server6"]["total"], 4)
self.assertEqual(errors["sources"]["server4"]["total"], 0)
self.assertEqual(
mod.workspace_state_summary(payload, "workspace unavailable"),
"workspace endpoint unavailable (using last good)",
)
fallback_error = copy.deepcopy(errors)
fallback_error["items"][0].pop("api_key_name", None)
fallback_error["items"][0].pop("account_name", None)
fallback_error["items"][0]["api_key_id"] = "9007199254740999"
fallback_error["items"][0]["account_id"] = "9007199254740998"
fallback_row = mod.normalize_error_rows(fallback_error)[0]
self.assertEqual(fallback_row["key"], "#9007199254740999")
self.assertEqual(fallback_row["account"], "#9007199254740998")
normalized_errors = mod.normalize_error_rows(errors)
self.assertEqual(normalized_errors[0]["count"], 4)
self.assertIn("count 4", mod.error_detail_line(normalized_errors[0]))
def test_workspace_payload_validation_rejects_nonfinite_and_oversized_rows(self) -> None:
mod = load_module()
nonfinite = workspace_payload_fixture()
nonfinite["traffic"]["requests"][0]["actual_cost"] = float("nan")
with mock.patch.object(mod, "fetch_payload", return_value=nonfinite):
with self.assertRaisesRegex(RuntimeError, "invalid workspace projection"):
mod.fetch_workspace_payload("https://workspace.example.test/data", 3)
oversized = workspace_payload_fixture()
oversized["traffic"]["requests"] = [
{"id": str(index)} for index in range(mod.MAX_WORKSPACE_REQUESTS + 1)
]
with mock.patch.object(mod, "fetch_payload", return_value=oversized):
with self.assertRaisesRegex(RuntimeError, "invalid workspace requests"):
mod.fetch_workspace_payload("https://workspace.example.test/data", 3)
def test_workspace_cache_coalesces_reads_and_returns_copies(self) -> None:
mod = load_module()
calls: list[str] = []
call_lock = threading.Lock()
def fake_request(url: str, timeout: int):
with call_lock:
calls.append(url)
time.sleep(0.03)
return workspace_payload_fixture()
cache = mod.WorkspaceCache("https://workspace.example.test/api/ui-data?view=workspace", 3, 300)
results: list[dict[str, object]] = []
threads = [
threading.Thread(target=lambda: results.append(cache.get()))
for _ in range(12)
]
with mock.patch.object(mod, "fetch_workspace_payload", side_effect=fake_request):
for thread in threads:
thread.start()
for thread in threads:
thread.join(timeout=2)
self.assertEqual(len(calls), 1)
self.assertEqual(len(results), 12)
results[0]["view"] = "mutated"
self.assertEqual(cache.get()["view"], "workspace")
cache.get(force=True)
self.assertEqual(len(calls), 2)
def test_workspace_cache_backs_off_failures_with_and_without_last_good(self) -> None:
mod = load_module()
cache = mod.WorkspaceCache("https://workspace.example.test/data", 3, 300)
with mock.patch.object(
mod,
"fetch_workspace_payload",
side_effect=[
workspace_payload_fixture(),
RuntimeError("private failure"),
workspace_payload_fixture(),
],
) as fetch:
first = cache.get()
stale = cache.get(force=True)
repeated = cache.get()
self.assertEqual(fetch.call_count, 2)
cache.last_attempt_at -= cache.retry_seconds + 1
recovered = cache.get()
self.assertEqual(first["view"], "workspace")
self.assertEqual(stale["view"], "workspace")
self.assertEqual(repeated["view"], "workspace")
self.assertEqual(recovered["view"], "workspace")
self.assertEqual(fetch.call_count, 3)
self.assertEqual(cache.network_fetches, 3)
self.assertEqual(cache.error, "")
empty = mod.WorkspaceCache("https://workspace.example.test/data", 3, 300)
with mock.patch.object(mod, "fetch_workspace_payload", side_effect=RuntimeError("private failure")) as fetch:
with self.assertRaisesRegex(RuntimeError, "workspace unavailable"):
empty.get()
with self.assertRaisesRegex(RuntimeError, "workspace unavailable"):
empty.get()
fetch.assert_called_once()
def test_default_main_uses_one_workspace_request_without_loading_admin_token(self) -> None:
mod = load_module()
out = io.StringIO()
with (
mock.patch.object(mod, "fetch_workspace_payload", return_value=workspace_payload_fixture()) as workspace_fetch,
mock.patch.object(mod, "default_logs_token", side_effect=AssertionError("admin token must not be read")) as token_read,
mock.patch.object(mod, "fetch_payload", side_effect=AssertionError("legacy accounts must not be read")),
mock.patch.object(mod, "fetch_optional_payload", side_effect=AssertionError("legacy status must not be read")),
mock.patch.object(mod, "fetch_optional_pricing_payload", side_effect=AssertionError("legacy pricing must not be read")),
contextlib.redirect_stdout(out),
):
rc = mod.main(
[
"--once",
"--workspace-url",
"https://workspace.example.test/api/ui-data?view=workspace",
"--no-version-check",
]
)
self.assertEqual(rc, 0)
workspace_fetch.assert_called_once_with(
"https://workspace.example.test/api/ui-data?view=workspace", 10
)
token_read.assert_not_called()
self.assertIn("oai-quota-code-plan", out.getvalue())
self.assertIn("wmy", out.getvalue())
def test_empty_workspace_url_falls_back_to_legacy_direct_mode(self) -> None:
mod = load_module()
accounts = workspace_payload_fixture()["accounts"]
with (
mock.patch.object(mod, "default_logs_token", return_value="") as token_read,
mock.patch.object(mod, "fetch_payload", return_value=accounts) as account_fetch,
mock.patch.object(mod, "fetch_optional_payload", return_value=({}, "")),
mock.patch.object(mod, "fetch_optional_pricing_payload", return_value=({"sources": []}, "")),
mock.patch.object(mod, "fetch_workspace_payload", side_effect=AssertionError("workspace must not be read")),
contextlib.redirect_stdout(io.StringIO()),
):
rc = mod.main(["--once", "--workspace-url", "", "--no-version-check"])
self.assertEqual(rc, 0)
token_read.assert_called_once_with()
account_fetch.assert_called_once()
def test_sources_and_requests_flags_are_aliases(self) -> None:
mod = load_module()
self.assertTrue(mod.build_parser().parse_args(["--sources"]).pricing)
self.assertTrue(mod.build_parser().parse_args(["--requests"]).logs)
class Sub2APILogsTests(unittest.TestCase):
def test_normalize_log_rows_maps_columns_and_sorts_newest_first(self) -> None:
mod = load_module()
@@ -737,7 +1049,7 @@ class Sub2APILogsTests(unittest.TestCase):
err = io.StringIO()
try:
with contextlib.redirect_stderr(err):
rc = mod.main(["--once", "--logs", "--no-version-check"])
rc = mod.main(["--once", "--logs", "--legacy-direct", "--no-version-check"])
finally:
if old_token is not None:
os.environ["SHUSUB2_LOGS_TOKEN"] = old_token
@@ -915,7 +1227,7 @@ class Sub2APILogsTests(unittest.TestCase):
err = io.StringIO()
try:
with contextlib.redirect_stderr(err):
rc = mod.main(["--once", "--errors", "--no-version-check"])
rc = mod.main(["--once", "--errors", "--legacy-direct", "--no-version-check"])
finally:
if old_token is not None:
os.environ["SHUSUB2_LOGS_TOKEN"] = old_token
@@ -1022,6 +1334,7 @@ class DashboardLayoutTests(unittest.IsolatedAsyncioTestCase):
100,
100,
"24h",
legacy_direct=True,
)
finally:
App.run = original_run
@@ -1061,6 +1374,102 @@ class DashboardLayoutTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(app.screen.rows[0]["name"], "code-plan")
async def test_workspace_dashboard_handles_initial_endpoint_failure(self) -> None:
from textual.app import App
mod = load_module()
captured: dict[str, object] = {}
original_run = App.run
App.run = lambda self, *args, **kwargs: captured.setdefault("app", self)
try:
rc = mod.run_textual(
"",
"",
"",
"",
"",
"",
"",
999,
999,
999,
1,
100,
100,
"24h",
workspace_url="https://workspace.example.test/api/ui-data?view=workspace",
legacy_direct=False,
)
finally:
App.run = original_run
self.assertEqual(rc, 0)
app = captured["app"]
with mock.patch.object(
mod,
"fetch_workspace_payload",
side_effect=RuntimeError("private endpoint failure"),
) as fetch:
async with app.run_test(size=(100, 30)) as pilot:
await pilot.pause()
await pilot.pause()
self.assertEqual(type(app.screen).__name__, "DashboardScreen")
self.assertIn("workspace unavailable", str(app.screen.query_one("#status").render()))
self.assertEqual(fetch.call_count, 1)
async def test_workspace_dashboard_uses_one_snapshot_and_has_return_navigation(self) -> None:
from textual.app import App
mod = load_module()
captured: dict[str, object] = {}
original_run = App.run
App.run = lambda self, *args, **kwargs: captured.setdefault("app", self)
try:
rc = mod.run_textual(
"",
"",
"",
"",
"",
"",
"",
999,
999,
999,
1,
100,
100,
"24h",
workspace_url="https://workspace.example.test/api/ui-data?view=workspace",
legacy_direct=False,
)
finally:
App.run = original_run
self.assertEqual(rc, 0)
app = captured["app"]
with mock.patch.object(mod, "fetch_workspace_payload", return_value=workspace_payload_fixture()) as fetch:
async with app.run_test(size=(100, 30)) as pilot:
await pilot.pause()
await pilot.pause()
screen = app.screen
self.assertEqual(type(screen).__name__, "DashboardScreen")
self.assertEqual(app.sub_title, "Dashboard")
self.assertEqual(screen.account_rows[0]["name"], "oai-quota-code-plan")
self.assertEqual(screen.error_rows[0]["count"], 4)
self.assertEqual(len(fetch.call_args_list), 1)
await pilot.press("p")
await pilot.pause()
self.assertEqual(type(app.screen).__name__, "PricingScreen")
self.assertEqual(app.sub_title, "Sources")
await pilot.press("d")
await pilot.pause()
self.assertEqual(type(app.screen).__name__, "DashboardScreen")
self.assertEqual(app.sub_title, "Dashboard")
self.assertEqual(len(fetch.call_args_list), 1)
class PageSelectionTests(unittest.TestCase):
def test_once_pricing_prints_pricing_monitor_sources(self) -> None:
mod = load_module()
@@ -1080,7 +1489,7 @@ class PageSelectionTests(unittest.TestCase):
out = io.StringIO()
try:
with contextlib.redirect_stdout(out):
rc = mod.main(["--once", "--pricing", "--no-version-check"])
rc = mod.main(["--once", "--pricing", "--legacy-direct", "--no-version-check"])
finally:
mod.fetch_pricing_payload = original_fetch
@@ -1097,12 +1506,12 @@ class PageSelectionTests(unittest.TestCase):
err = io.StringIO()
try:
with contextlib.redirect_stderr(err):
rc = mod.main(["--once", "--pricing", "--no-version-check"])
rc = mod.main(["--once", "--pricing", "--legacy-direct", "--no-version-check"])
finally:
mod.fetch_pricing_payload = original_fetch
self.assertEqual(rc, 1)
self.assertEqual(err.getvalue().strip(), "upstreams unavailable")
self.assertEqual(err.getvalue().strip(), "sources unavailable")
def test_once_accounts_includes_only_mapped_pricing_source_balance(self) -> None:
mod = load_module()
@@ -1136,6 +1545,7 @@ class PageSelectionTests(unittest.TestCase):
with contextlib.redirect_stdout(out):
rc = mod.main(
[
"--legacy-direct",
"--once",
"--api-url",
"https://accounts.example.test/api/tui/accounts",