422 lines
16 KiB
Python
422 lines
16 KiB
Python
#!/usr/bin/env python3
|
|
"""Export first-party coding CLI quotas without exposing account credentials."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import os
|
|
import selectors
|
|
import subprocess
|
|
import threading
|
|
import time
|
|
from dataclasses import dataclass, field
|
|
from datetime import UTC, date, datetime, timedelta
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
from pathlib import Path
|
|
from typing import Any
|
|
from urllib.request import Request, urlopen
|
|
|
|
|
|
CODEX_BIN = os.environ.get("ATLAS_AI_CODEX_BIN", "/opt/data/tools/bin/codex")
|
|
CODEX_HOME = os.environ.get("CODEX_HOME", "/runtime-access/codex")
|
|
CLAUDE_CREDENTIALS = Path(
|
|
os.environ.get(
|
|
"ATLAS_AI_CLAUDE_CREDENTIALS",
|
|
"/runtime-access/claude/.credentials.json",
|
|
)
|
|
)
|
|
CLAUDE_USAGE_URL = os.environ.get(
|
|
"ATLAS_AI_CLAUDE_USAGE_URL",
|
|
"https://api.anthropic.com/api/oauth/usage",
|
|
)
|
|
CLAUDE_WINDOWS = (
|
|
"five_hour",
|
|
"seven_day",
|
|
"seven_day_opus",
|
|
"seven_day_sonnet",
|
|
)
|
|
METRIC_HELP = {
|
|
"atlas_ai_account_tokens": "First-party account token usage for a fixed period.",
|
|
"atlas_ai_account_usage_summary": "First-party account usage summary values.",
|
|
"atlas_ai_extra_usage_enabled": "Whether metered extra usage is enabled for the account.",
|
|
"atlas_ai_quota_fetch_duration_seconds": "Duration of the latest provider quota fetch.",
|
|
"atlas_ai_quota_fetch_success": "Whether the latest provider quota fetch succeeded.",
|
|
"atlas_ai_quota_last_attempt_timestamp_seconds": "Unix timestamp of the latest quota fetch attempt.",
|
|
"atlas_ai_quota_last_success_timestamp_seconds": "Unix timestamp of the latest successful quota fetch.",
|
|
"atlas_ai_quota_remaining_percent": "Remaining percentage in a first-party coding CLI quota window.",
|
|
"atlas_ai_quota_reset_timestamp_seconds": "Unix timestamp when a coding CLI quota window resets.",
|
|
"atlas_ai_quota_used_percent": "Used percentage in a first-party coding CLI quota window.",
|
|
"atlas_ai_quota_window_duration_seconds": "Nominal duration of a coding CLI quota window.",
|
|
}
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Sample:
|
|
"""One Prometheus gauge sample."""
|
|
|
|
name: str
|
|
labels: dict[str, str]
|
|
value: float
|
|
|
|
|
|
@dataclass
|
|
class ProviderState:
|
|
"""Latest safe samples and fetch health for one provider."""
|
|
|
|
samples: list[Sample] = field(default_factory=list)
|
|
fetch_success: bool = False
|
|
last_attempt: float = 0
|
|
last_success: float = 0
|
|
duration: float = 0
|
|
|
|
|
|
def _number(value: Any) -> float | None:
|
|
"""Return a finite numeric value while rejecting booleans and nulls."""
|
|
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
|
return None
|
|
converted = float(value)
|
|
return converted if converted == converted and abs(converted) != float("inf") else None
|
|
|
|
|
|
def _timestamp(value: Any) -> float | None:
|
|
"""Parse either a Unix timestamp or an ISO-8601 timestamp."""
|
|
number = _number(value)
|
|
if number is not None:
|
|
return number
|
|
if not isinstance(value, str) or not value:
|
|
return None
|
|
try:
|
|
return datetime.fromisoformat(value.replace("Z", "+00:00")).timestamp()
|
|
except ValueError:
|
|
return None
|
|
|
|
|
|
def _window_name(window: dict[str, Any], fallback: str) -> str:
|
|
"""Give common rolling windows stable, human-readable labels."""
|
|
minutes = _number(window.get("windowDurationMins"))
|
|
known = {300: "five_hour", 1440: "one_day", 10080: "seven_day"}
|
|
if minutes is not None and int(minutes) in known:
|
|
return known[int(minutes)]
|
|
return fallback
|
|
|
|
|
|
def _codex_limit_name(limit_id: str, snapshot: dict[str, Any]) -> str:
|
|
"""Return a stable low-cardinality name for a Codex quota bucket."""
|
|
if limit_id == "codex":
|
|
return "codex"
|
|
name = snapshot.get("limitName")
|
|
if isinstance(name, str) and name:
|
|
normalized = "".join(char.lower() if char.isalnum() else "-" for char in name)
|
|
return "-".join(filter(None, normalized.split("-")))[:64]
|
|
return "additional"
|
|
|
|
|
|
def parse_codex_payloads(
|
|
rate_response: dict[str, Any], usage_response: dict[str, Any], *, today: date | None = None
|
|
) -> list[Sample]:
|
|
"""Convert structured Codex app-server responses into bounded metrics."""
|
|
samples: list[Sample] = []
|
|
limits = rate_response.get("rateLimitsByLimitId")
|
|
if not isinstance(limits, dict) or not limits:
|
|
limits = {"codex": rate_response.get("rateLimits")}
|
|
for limit_id, snapshot in limits.items():
|
|
if not isinstance(limit_id, str) or not isinstance(snapshot, dict):
|
|
continue
|
|
limit_name = _codex_limit_name(limit_id, snapshot)
|
|
for fallback, raw_window in (
|
|
("primary", snapshot.get("primary")),
|
|
("secondary", snapshot.get("secondary")),
|
|
):
|
|
if not isinstance(raw_window, dict):
|
|
continue
|
|
used = _number(raw_window.get("usedPercent"))
|
|
if used is None:
|
|
continue
|
|
labels = {
|
|
"provider": "openai",
|
|
"limit": limit_name,
|
|
"window": _window_name(raw_window, fallback),
|
|
}
|
|
samples.extend(
|
|
(
|
|
Sample("atlas_ai_quota_used_percent", labels, used),
|
|
Sample("atlas_ai_quota_remaining_percent", labels, max(0, 100 - used)),
|
|
)
|
|
)
|
|
reset = _timestamp(raw_window.get("resetsAt"))
|
|
duration = _number(raw_window.get("windowDurationMins"))
|
|
if reset is not None:
|
|
samples.append(Sample("atlas_ai_quota_reset_timestamp_seconds", labels, reset))
|
|
if duration is not None:
|
|
samples.append(Sample("atlas_ai_quota_window_duration_seconds", labels, duration * 60))
|
|
|
|
summary = usage_response.get("summary")
|
|
if isinstance(summary, dict):
|
|
for source, metric in (
|
|
("lifetimeTokens", "lifetime_tokens"),
|
|
("peakDailyTokens", "peak_daily_tokens"),
|
|
("currentStreakDays", "current_streak_days"),
|
|
("longestStreakDays", "longest_streak_days"),
|
|
("longestRunningTurnSec", "longest_running_turn_seconds"),
|
|
):
|
|
value = _number(summary.get(source))
|
|
if value is not None:
|
|
samples.append(
|
|
Sample(
|
|
"atlas_ai_account_usage_summary",
|
|
{"provider": "openai", "metric": metric},
|
|
value,
|
|
)
|
|
)
|
|
|
|
current_day = today or datetime.now(UTC).date()
|
|
daily = usage_response.get("dailyUsageBuckets")
|
|
parsed_daily: list[tuple[date, float]] = []
|
|
if isinstance(daily, list):
|
|
for bucket in daily:
|
|
if not isinstance(bucket, dict):
|
|
continue
|
|
value = _number(bucket.get("tokens"))
|
|
try:
|
|
start = date.fromisoformat(str(bucket.get("startDate")))
|
|
except ValueError:
|
|
continue
|
|
if value is not None:
|
|
parsed_daily.append((start, value))
|
|
# The account endpoint publishes completed daily buckets and may not include
|
|
# the current UTC day. Anchor fixed periods to the latest reported day so a
|
|
# delayed bucket is not mislabeled as zero usage.
|
|
anchor_day = max((start for start, _ in parsed_daily), default=current_day)
|
|
for period, days in (("latest_day", 1), ("seven_day", 7), ("thirty_day", 30)):
|
|
earliest = anchor_day - timedelta(days=days - 1)
|
|
value = sum(tokens for start, tokens in parsed_daily if earliest <= start <= anchor_day)
|
|
samples.append(Sample("atlas_ai_account_tokens", {"provider": "openai", "period": period}, value))
|
|
return samples
|
|
|
|
|
|
def parse_claude_payload(payload: dict[str, Any]) -> list[Sample]:
|
|
"""Convert Claude's first-party OAuth usage document into bounded metrics."""
|
|
samples: list[Sample] = []
|
|
for window_name in CLAUDE_WINDOWS:
|
|
window = payload.get(window_name)
|
|
if not isinstance(window, dict):
|
|
continue
|
|
used = _number(window.get("utilization"))
|
|
if used is None:
|
|
continue
|
|
labels = {"provider": "anthropic", "limit": "claude", "window": window_name}
|
|
samples.extend(
|
|
(
|
|
Sample("atlas_ai_quota_used_percent", labels, used),
|
|
Sample("atlas_ai_quota_remaining_percent", labels, max(0, 100 - used)),
|
|
)
|
|
)
|
|
reset = _timestamp(window.get("resets_at"))
|
|
if reset is not None:
|
|
samples.append(Sample("atlas_ai_quota_reset_timestamp_seconds", labels, reset))
|
|
extra = payload.get("extra_usage")
|
|
if isinstance(extra, dict):
|
|
enabled = extra.get("is_enabled")
|
|
if isinstance(enabled, bool):
|
|
samples.append(
|
|
Sample(
|
|
"atlas_ai_extra_usage_enabled",
|
|
{"provider": "anthropic"},
|
|
float(enabled),
|
|
)
|
|
)
|
|
return samples
|
|
|
|
|
|
def query_codex(timeout: float = 20) -> tuple[dict[str, Any], dict[str, Any]]:
|
|
"""Read Codex account quota and usage through its structured app-server protocol."""
|
|
process = subprocess.Popen(
|
|
[CODEX_BIN, "app-server", "--stdio"],
|
|
stdin=subprocess.PIPE,
|
|
stdout=subprocess.PIPE,
|
|
stderr=subprocess.DEVNULL,
|
|
text=True,
|
|
bufsize=1,
|
|
env={**os.environ, "CODEX_HOME": CODEX_HOME},
|
|
)
|
|
requests = (
|
|
{
|
|
"id": 1,
|
|
"method": "initialize",
|
|
"params": {
|
|
"clientInfo": {
|
|
"name": "atlas-ai-usage-exporter",
|
|
"title": "Atlas AI Usage Exporter",
|
|
"version": "1.0.0",
|
|
},
|
|
"capabilities": {"experimentalApi": True},
|
|
},
|
|
},
|
|
{"method": "initialized", "params": {}},
|
|
{"id": 2, "method": "account/rateLimits/read", "params": None},
|
|
{"id": 3, "method": "account/usage/read", "params": None},
|
|
)
|
|
try:
|
|
if process.stdin is None or process.stdout is None:
|
|
raise RuntimeError("Codex app-server pipes are unavailable")
|
|
for message in requests:
|
|
process.stdin.write(json.dumps(message, separators=(",", ":")) + "\n")
|
|
process.stdin.flush()
|
|
selector = selectors.DefaultSelector()
|
|
selector.register(process.stdout, selectors.EVENT_READ)
|
|
responses: dict[int, dict[str, Any]] = {}
|
|
deadline = time.monotonic() + timeout
|
|
while len(responses) < 3 and time.monotonic() < deadline:
|
|
for key, _ in selector.select(min(1, max(0, deadline - time.monotonic()))):
|
|
line = key.fileobj.readline()
|
|
if not line:
|
|
continue
|
|
message = json.loads(line)
|
|
if message.get("id") in (1, 2, 3):
|
|
responses[int(message["id"])] = message
|
|
for response_id in (1, 2, 3):
|
|
response = responses.get(response_id)
|
|
if not response or "error" in response or not isinstance(response.get("result"), dict):
|
|
raise RuntimeError(f"Codex app-server response {response_id} failed")
|
|
return responses[2]["result"], responses[3]["result"]
|
|
finally:
|
|
process.terminate()
|
|
try:
|
|
process.wait(timeout=5)
|
|
except subprocess.TimeoutExpired:
|
|
process.kill()
|
|
process.wait(timeout=5)
|
|
|
|
|
|
def query_claude() -> dict[str, Any]:
|
|
"""Read Claude account quota with the runtime OAuth token held only in memory."""
|
|
document = json.loads(CLAUDE_CREDENTIALS.read_text(encoding="utf-8"))
|
|
token = document.get("claudeAiOauth", {}).get("accessToken")
|
|
if not isinstance(token, str) or not token:
|
|
raise RuntimeError("Claude runtime credentials are incomplete")
|
|
request = Request(
|
|
CLAUDE_USAGE_URL,
|
|
headers={
|
|
"Authorization": f"Bearer {token}",
|
|
"Accept": "application/json",
|
|
"anthropic-version": "2023-06-01",
|
|
"User-Agent": "atlas-ai-usage-exporter/1.0",
|
|
},
|
|
)
|
|
with urlopen(request, timeout=15) as response:
|
|
payload = json.loads(response.read(1 << 20))
|
|
if not isinstance(payload, dict):
|
|
raise RuntimeError("Claude usage response is not an object")
|
|
return payload
|
|
|
|
|
|
def _escape_label(value: str) -> str:
|
|
"""Escape one Prometheus label value."""
|
|
return value.replace("\\", "\\\\").replace("\n", "\\n").replace('"', '\\"')
|
|
|
|
|
|
class Collector:
|
|
"""Poll providers independently and expose a thread-safe metrics snapshot."""
|
|
|
|
def __init__(self) -> None:
|
|
self._lock = threading.Lock()
|
|
self._providers = {name: ProviderState() for name in ("openai", "anthropic")}
|
|
|
|
def refresh_provider(self, provider: str) -> None:
|
|
"""Refresh one provider while retaining the last good values on failure."""
|
|
started = time.time()
|
|
monotonic_started = time.monotonic()
|
|
try:
|
|
if provider == "openai":
|
|
samples = parse_codex_payloads(*query_codex())
|
|
elif provider == "anthropic":
|
|
samples = parse_claude_payload(query_claude())
|
|
else:
|
|
raise ValueError("unknown provider")
|
|
success = True
|
|
except (OSError, RuntimeError, ValueError, json.JSONDecodeError) as error:
|
|
print(f"{provider} quota collection deferred: {type(error).__name__}", flush=True)
|
|
samples = []
|
|
success = False
|
|
with self._lock:
|
|
state = self._providers[provider]
|
|
state.last_attempt = started
|
|
state.duration = time.monotonic() - monotonic_started
|
|
state.fetch_success = success
|
|
if success:
|
|
state.samples = samples
|
|
state.last_success = time.time()
|
|
|
|
def render(self) -> bytes:
|
|
"""Render the current provider states in Prometheus text format."""
|
|
with self._lock:
|
|
states = {
|
|
provider: ProviderState(**vars(state))
|
|
for provider, state in self._providers.items()
|
|
}
|
|
samples: list[Sample] = []
|
|
for provider, state in states.items():
|
|
labels = {"provider": provider}
|
|
samples.extend(state.samples)
|
|
samples.extend(
|
|
(
|
|
Sample("atlas_ai_quota_fetch_success", labels, float(state.fetch_success)),
|
|
Sample("atlas_ai_quota_last_attempt_timestamp_seconds", labels, state.last_attempt),
|
|
Sample("atlas_ai_quota_last_success_timestamp_seconds", labels, state.last_success),
|
|
Sample("atlas_ai_quota_fetch_duration_seconds", labels, state.duration),
|
|
)
|
|
)
|
|
lines: list[str] = []
|
|
for name in sorted({sample.name for sample in samples}):
|
|
lines.extend((f"# HELP {name} {METRIC_HELP[name]}", f"# TYPE {name} gauge"))
|
|
for sample in sorted(
|
|
(item for item in samples if item.name == name),
|
|
key=lambda item: sorted(item.labels.items()),
|
|
):
|
|
labels = ",".join(
|
|
f'{key}="{_escape_label(value)}"'
|
|
for key, value in sorted(sample.labels.items())
|
|
)
|
|
lines.append(f"{name}{{{labels}}} {sample.value:.12g}")
|
|
return ("\n".join(lines) + "\n").encode("utf-8")
|
|
|
|
|
|
def main() -> int:
|
|
"""Poll quota APIs and serve only sanitized metrics and health endpoints."""
|
|
collector = Collector()
|
|
interval = max(60, int(os.environ.get("ATLAS_AI_USAGE_INTERVAL_SECONDS", "300")))
|
|
|
|
def polling_loop() -> None:
|
|
while True:
|
|
for provider in ("openai", "anthropic"):
|
|
collector.refresh_provider(provider)
|
|
time.sleep(interval)
|
|
|
|
class Handler(BaseHTTPRequestHandler):
|
|
def do_GET(self) -> None: # noqa: N802
|
|
if self.path == "/metrics":
|
|
payload = collector.render()
|
|
self.send_response(200)
|
|
self.send_header("Content-Type", "text/plain; version=0.0.4")
|
|
self.send_header("Content-Length", str(len(payload)))
|
|
self.end_headers()
|
|
self.wfile.write(payload)
|
|
elif self.path == "/healthz":
|
|
self.send_response(200)
|
|
self.end_headers()
|
|
else:
|
|
self.send_error(404)
|
|
|
|
def log_message(self, _format: str, *_args: object) -> None:
|
|
return
|
|
|
|
threading.Thread(target=polling_loop, daemon=True).start()
|
|
port = int(os.environ.get("ATLAS_AI_USAGE_PORT", "9010"))
|
|
server = ThreadingHTTPServer(("0.0.0.0", port), Handler)
|
|
server.serve_forever()
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|