chopratejas commited on
Commit
cf08723
·
1 Parent(s): 07c6b43

feat: add license key validation + phone-home usage reporter

Browse files

Enables managed/enterprise customers to run the proxy in their own
environment while reporting aggregate usage back to Headroom cloud.

- HEADROOM_LICENSE_KEY env var activates managed mode
- UsageReporter validates license on startup, caches to ~/.headroom/
- Reports aggregate stats every 5 min (tokens, costs, models — no content)
- Graceful degradation: 7-day grace period if cloud unreachable
- Expired license → passthrough mode (proxy works, compression stops)
- Zero impact on OSS users (no license key = no reporter)

headroom/cli/proxy.py CHANGED
@@ -201,6 +201,9 @@ def proxy(
201
  # Resolve mode: CLI flag > env var > default
202
  effective_mode = mode or os.environ.get("HEADROOM_MODE", "token_headroom")
203
 
 
 
 
204
  config = ProxyConfig(
205
  host=host,
206
  port=port,
@@ -238,12 +241,18 @@ def proxy(
238
  bedrock_region=bedrock_region or region,
239
  bedrock_profile=bedrock_profile,
240
  anyllm_provider=effective_anyllm_provider,
 
 
241
  )
242
 
243
  memory_status = "DISABLED"
244
  if config.memory_enabled:
245
  memory_status = "ENABLED (multi-provider)"
246
 
 
 
 
 
247
  effective_region = bedrock_region or region
248
  backend_status = "Anthropic (direct API)"
249
  backend_section = ""
@@ -314,6 +323,7 @@ Starting proxy server...
314
  Caching: {"ENABLED" if config.cache_enabled else "DISABLED"}
315
  Rate Limit: {"ENABLED" if config.rate_limit_enabled else "DISABLED"}
316
  Memory: {memory_status}
 
317
  {backend_section}
318
  Usage with Claude Code:
319
  ANTHROPIC_BASE_URL=http://{config.host}:{config.port} claude
 
201
  # Resolve mode: CLI flag > env var > default
202
  effective_mode = mode or os.environ.get("HEADROOM_MODE", "token_headroom")
203
 
204
+ # License key for managed/enterprise deployments (optional)
205
+ license_key = os.environ.get("HEADROOM_LICENSE_KEY")
206
+
207
  config = ProxyConfig(
208
  host=host,
209
  port=port,
 
241
  bedrock_region=bedrock_region or region,
242
  bedrock_profile=bedrock_profile,
243
  anyllm_provider=effective_anyllm_provider,
244
+ # License / Usage Reporting (managed/enterprise)
245
+ license_key=license_key,
246
  )
247
 
248
  memory_status = "DISABLED"
249
  if config.memory_enabled:
250
  memory_status = "ENABLED (multi-provider)"
251
 
252
+ license_status = "OSS (no license key)"
253
+ if license_key:
254
+ license_status = f"MANAGED (key={license_key[:8]}...)"
255
+
256
  effective_region = bedrock_region or region
257
  backend_status = "Anthropic (direct API)"
258
  backend_section = ""
 
323
  Caching: {"ENABLED" if config.cache_enabled else "DISABLED"}
324
  Rate Limit: {"ENABLED" if config.rate_limit_enabled else "DISABLED"}
325
  Memory: {memory_status}
326
+ License: {license_status}
327
  {backend_section}
328
  Usage with Claude Code:
329
  ANTHROPIC_BASE_URL=http://{config.host}:{config.port} claude
headroom/proxy/server.py CHANGED
@@ -733,6 +733,11 @@ class ProxyConfig:
733
  memory_bridge_auto_import: bool = False
734
  memory_bridge_export_path: str = ""
735
 
 
 
 
 
 
736
  # Compression Hooks (for SaaS and advanced customization)
737
  hooks: Any = None # CompressionHooks instance, or None for default behavior
738
 
@@ -1842,6 +1847,17 @@ class HeadroomProxy:
1842
  )
1843
  self.memory_handler = MemoryHandler(memory_config)
1844
 
 
 
 
 
 
 
 
 
 
 
 
1845
  # Traffic Learner (live pattern extraction from proxy traffic)
1846
  # Only activates with --learn flag; requires --memory for backend
1847
  self.traffic_learner: TrafficLearner | None = None
@@ -2371,7 +2387,8 @@ class HeadroomProxy:
2371
 
2372
  _compression_failed = False
2373
  original_messages = messages # Preserve for 400-retry fallback
2374
- if self.config.optimize and messages and not _bypass:
 
2375
  try:
2376
  context_limit = self.anthropic_provider.get_context_limit(model)
2377
  biases = (
@@ -4562,8 +4579,12 @@ class HeadroomProxy:
4562
  -MAX_SSE_BUFFER_SIZE // 2 :
4563
  ]
4564
 
 
 
 
 
4565
  if memory_enabled:
4566
- # Buffer for memory tool detection
4567
  buffered_chunks.append(chunk)
4568
  full_sse_data += chunk_str
4569
  if len(full_sse_data) > MAX_SSE_BUFFER_SIZE:
@@ -4572,9 +4593,6 @@ class HeadroomProxy:
4572
  "disabling memory detection for this request"
4573
  )
4574
  memory_enabled = False
4575
- else:
4576
- # Immediate streaming when memory not enabled
4577
- yield chunk
4578
 
4579
  # Parse complete SSE events from buffer
4580
  usage = self._parse_sse_usage_from_buffer(stream_state, provider)
@@ -4593,17 +4611,16 @@ class HeadroomProxy:
4593
  ]
4594
 
4595
  # Memory tool handling after stream completes
 
 
4596
  if memory_enabled and full_sse_data:
4597
- # Check for Claude Code credential error in initial response
4598
  if "only authorized for use with Claude Code" in full_sse_data:
4599
  logger.warning(
4600
  f"[{request_id}] Memory: Claude Code subscription credentials "
4601
  "do not support custom tool injection. Set ANTHROPIC_API_KEY "
4602
  "environment variable or use --no-memory-tools flag."
4603
  )
4604
- # Yield buffered error response as-is (contains error details)
4605
- for chunk in buffered_chunks:
4606
- yield chunk
4607
  return
4608
 
4609
  # Parse SSE to get response JSON
@@ -4616,78 +4633,16 @@ class HeadroomProxy:
4616
  f"[{request_id}] Memory: Detected tool calls in streaming response"
4617
  )
4618
 
4619
- # Execute memory tool calls
 
4620
  tool_results = await self.memory_handler.handle_memory_tool_calls(
4621
  parsed_response, memory_user_id, provider
4622
  )
4623
-
4624
  if tool_results:
4625
- # Build continuation messages
4626
- # Filter out system role messages (Anthropic uses top-level 'system' param)
4627
- messages = [
4628
- m for m in body.get("messages", []) if m.get("role") != "system"
4629
- ]
4630
- assistant_msg = {
4631
- "role": "assistant",
4632
- "content": parsed_response.get("content", []),
4633
- }
4634
- user_msg = {"role": "user", "content": tool_results}
4635
- continuation_messages = messages + [assistant_msg, user_msg]
4636
-
4637
- # Make continuation request (streaming to support Claude Code API key)
4638
- continuation_body = {
4639
- **body,
4640
- "messages": continuation_messages,
4641
- "stream": True,
4642
- }
4643
-
4644
  logger.info(
4645
- f"[{request_id}] Memory: Tool execution complete, streaming continuation"
 
4646
  )
4647
-
4648
- # Stream continuation response directly
4649
- async with self.http_client.stream(
4650
- "POST", url, json=continuation_body, headers=headers
4651
- ) as cont_response:
4652
- cont_buffer = ""
4653
- async for chunk in cont_response.aiter_bytes():
4654
- chunk_str = chunk.decode("utf-8", errors="ignore")
4655
- cont_buffer += chunk_str
4656
-
4657
- # Check for Claude Code credential error
4658
- if "only authorized for use with Claude Code" in cont_buffer:
4659
- logger.warning(
4660
- f"[{request_id}] Memory: Claude Code subscription "
4661
- "credentials do not support custom tool injection. "
4662
- "Set ANTHROPIC_API_KEY environment variable to use "
4663
- "memory tools, or disable memory tools with "
4664
- "--no-memory-tools flag."
4665
- )
4666
- # Yield a helpful error message in SSE format
4667
- error_msg = (
4668
- "Memory tools require a regular Anthropic API key. "
4669
- "Claude Code subscription credentials do not allow "
4670
- "custom tool injection. "
4671
- "To fix: (1) Set ANTHROPIC_API_KEY=your_api_key before "
4672
- "starting the proxy, or (2) Run proxy with "
4673
- "--no-memory-tools flag."
4674
- )
4675
- error_event = (
4676
- f'data: {{"type":"error","error":{{"type":"permission_error",'
4677
- f'"message":"{error_msg}"}}}}\n\n'
4678
- )
4679
- yield error_event.encode()
4680
- return
4681
-
4682
- yield chunk
4683
- else:
4684
- # No tool results, yield original buffered chunks
4685
- for chunk in buffered_chunks:
4686
- yield chunk
4687
- else:
4688
- # No memory tool calls, yield original buffered chunks
4689
- for chunk in buffered_chunks:
4690
- yield chunk
4691
  except (httpx.ConnectError, httpx.ConnectTimeout, httpx.PoolTimeout) as e:
4692
  logger.error(f"[{request_id}] Connection error to upstream API: {e}")
4693
  error_event = {
@@ -5124,7 +5079,8 @@ class HeadroomProxy:
5124
 
5125
  _compression_failed = False
5126
  original_messages = messages # Preserve for 400-retry fallback
5127
- if self.config.optimize and messages and not _bypass:
 
5128
  try:
5129
  context_limit = self.openai_provider.get_context_limit(model)
5130
 
@@ -6388,7 +6344,8 @@ class HeadroomProxy:
6388
  optimized_tokens = original_tokens
6389
 
6390
  _compression_failed = False
6391
- if self.config.optimize and messages:
 
6392
  try:
6393
  # Use OpenAI pipeline (similar message format)
6394
  context_limit = self.openai_provider.get_context_limit(model)
@@ -6880,12 +6837,17 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI:
6880
  await proxy.startup()
6881
  # Start background task for periodic TOIN stats logging
6882
  asyncio.create_task(_log_toin_stats_periodically())
 
 
 
6883
  # Start traffic learner background save worker
6884
  if proxy.traffic_learner:
6885
  await proxy.traffic_learner.start()
6886
 
6887
  @app.on_event("shutdown")
6888
  async def shutdown():
 
 
6889
  if proxy.traffic_learner:
6890
  await proxy.traffic_learner.stop()
6891
  await proxy.shutdown()
 
733
  memory_bridge_auto_import: bool = False
734
  memory_bridge_export_path: str = ""
735
 
736
+ # License / Usage Reporting (managed/enterprise deployments)
737
+ license_key: str | None = None # HEADROOM_LICENSE_KEY env var
738
+ license_cloud_url: str = "https://app.headroomlabs.ai"
739
+ license_report_interval: int = 300 # seconds (5 min)
740
+
741
  # Compression Hooks (for SaaS and advanced customization)
742
  hooks: Any = None # CompressionHooks instance, or None for default behavior
743
 
 
1847
  )
1848
  self.memory_handler = MemoryHandler(memory_config)
1849
 
1850
+ # Usage Reporter (license validation + phone-home for managed/enterprise)
1851
+ self.usage_reporter: UsageReporter | None = None
1852
+ if config.license_key:
1853
+ from headroom.telemetry.reporter import UsageReporter
1854
+
1855
+ self.usage_reporter = UsageReporter(
1856
+ license_key=config.license_key,
1857
+ cloud_url=config.license_cloud_url,
1858
+ report_interval=config.license_report_interval,
1859
+ )
1860
+
1861
  # Traffic Learner (live pattern extraction from proxy traffic)
1862
  # Only activates with --learn flag; requires --memory for backend
1863
  self.traffic_learner: TrafficLearner | None = None
 
2387
 
2388
  _compression_failed = False
2389
  original_messages = messages # Preserve for 400-retry fallback
2390
+ _license_ok = self.usage_reporter.should_compress if self.usage_reporter else True
2391
+ if self.config.optimize and messages and not _bypass and _license_ok:
2392
  try:
2393
  context_limit = self.anthropic_provider.get_context_limit(model)
2394
  biases = (
 
4579
  -MAX_SSE_BUFFER_SIZE // 2 :
4580
  ]
4581
 
4582
+ # Always stream immediately — buffering breaks
4583
+ # real-time clients (LangGraph, LangChain, etc.)
4584
+ yield chunk
4585
+
4586
  if memory_enabled:
4587
+ # Also buffer for post-stream memory processing
4588
  buffered_chunks.append(chunk)
4589
  full_sse_data += chunk_str
4590
  if len(full_sse_data) > MAX_SSE_BUFFER_SIZE:
 
4593
  "disabling memory detection for this request"
4594
  )
4595
  memory_enabled = False
 
 
 
4596
 
4597
  # Parse complete SSE events from buffer
4598
  usage = self._parse_sse_usage_from_buffer(stream_state, provider)
 
4611
  ]
4612
 
4613
  # Memory tool handling after stream completes
4614
+ # Chunks were already yielded in real-time above, so we only
4615
+ # do silent background processing here — no yielding.
4616
  if memory_enabled and full_sse_data:
4617
+ # Check for Claude Code credential error
4618
  if "only authorized for use with Claude Code" in full_sse_data:
4619
  logger.warning(
4620
  f"[{request_id}] Memory: Claude Code subscription credentials "
4621
  "do not support custom tool injection. Set ANTHROPIC_API_KEY "
4622
  "environment variable or use --no-memory-tools flag."
4623
  )
 
 
 
4624
  return
4625
 
4626
  # Parse SSE to get response JSON
 
4633
  f"[{request_id}] Memory: Detected tool calls in streaming response"
4634
  )
4635
 
4636
+ # Execute memory tool calls silently — response already
4637
+ # streamed so we cannot make a continuation request.
4638
  tool_results = await self.memory_handler.handle_memory_tool_calls(
4639
  parsed_response, memory_user_id, provider
4640
  )
 
4641
  if tool_results:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4642
  logger.info(
4643
+ f"[{request_id}] Memory: Tool calls executed silently "
4644
+ "(streaming mode — no continuation)"
4645
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4646
  except (httpx.ConnectError, httpx.ConnectTimeout, httpx.PoolTimeout) as e:
4647
  logger.error(f"[{request_id}] Connection error to upstream API: {e}")
4648
  error_event = {
 
5079
 
5080
  _compression_failed = False
5081
  original_messages = messages # Preserve for 400-retry fallback
5082
+ _license_ok = self.usage_reporter.should_compress if self.usage_reporter else True
5083
+ if self.config.optimize and messages and not _bypass and _license_ok:
5084
  try:
5085
  context_limit = self.openai_provider.get_context_limit(model)
5086
 
 
6344
  optimized_tokens = original_tokens
6345
 
6346
  _compression_failed = False
6347
+ _license_ok = self.usage_reporter.should_compress if self.usage_reporter else True
6348
+ if self.config.optimize and messages and _license_ok:
6349
  try:
6350
  # Use OpenAI pipeline (similar message format)
6351
  context_limit = self.openai_provider.get_context_limit(model)
 
6837
  await proxy.startup()
6838
  # Start background task for periodic TOIN stats logging
6839
  asyncio.create_task(_log_toin_stats_periodically())
6840
+ # Start usage reporter (license validation + phone-home)
6841
+ if proxy.usage_reporter:
6842
+ await proxy.usage_reporter.start(proxy)
6843
  # Start traffic learner background save worker
6844
  if proxy.traffic_learner:
6845
  await proxy.traffic_learner.start()
6846
 
6847
  @app.on_event("shutdown")
6848
  async def shutdown():
6849
+ if proxy.usage_reporter:
6850
+ await proxy.usage_reporter.stop()
6851
  if proxy.traffic_learner:
6852
  await proxy.traffic_learner.stop()
6853
  await proxy.shutdown()
headroom/telemetry/reporter.py ADDED
@@ -0,0 +1,387 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """License validation and usage reporting for managed/enterprise deployments.
2
+
3
+ Phone-home module that validates license keys and reports aggregate usage
4
+ statistics to the Headroom cloud for billing. Designed to be non-intrusive:
5
+ the proxy works normally even if the cloud API is completely unreachable.
6
+
7
+ Privacy guarantees:
8
+ - Never sends message content, API keys, prompts, tool results, or user data
9
+ - Only sends aggregate counts: requests, tokens saved, model distribution
10
+ - All communication is over HTTPS
11
+
12
+ Usage:
13
+ reporter = UsageReporter(license_key="hlk_...")
14
+ license_info = await reporter.validate_license()
15
+ await reporter.start(proxy) # starts background loop
16
+ ...
17
+ await reporter.stop()
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ import asyncio
23
+ import json
24
+ import logging
25
+ from dataclasses import asdict, dataclass, field
26
+ from datetime import datetime, timezone
27
+ from pathlib import Path
28
+ from typing import TYPE_CHECKING, Any
29
+
30
+ import httpx
31
+
32
+ if TYPE_CHECKING:
33
+ from headroom.proxy.server import HeadroomProxy
34
+
35
+ logger = logging.getLogger("headroom.telemetry.reporter")
36
+
37
+ # Grace period: if the cloud API is unreachable, use cached license for up to 7 days
38
+ GRACE_PERIOD_SECONDS = 7 * 24 * 3600 # 7 days
39
+
40
+ # Default cache location
41
+ LICENSE_CACHE_PATH = Path.home() / ".headroom" / "license_cache.json"
42
+
43
+
44
+ @dataclass
45
+ class LicenseInfo:
46
+ """Cached license validation result."""
47
+
48
+ status: str # "active", "trial", "expired", "invalid"
49
+ org_id: str | None = None
50
+ org_name: str | None = None
51
+ plan: str | None = None
52
+ quota_tokens: int | None = None # None = unlimited
53
+ trial_expires_at: datetime | None = None
54
+ validated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
55
+
56
+ def to_dict(self) -> dict[str, Any]:
57
+ """Serialize for JSON caching."""
58
+ d = asdict(self)
59
+ d["validated_at"] = self.validated_at.isoformat()
60
+ if self.trial_expires_at:
61
+ d["trial_expires_at"] = self.trial_expires_at.isoformat()
62
+ return d
63
+
64
+ @classmethod
65
+ def from_dict(cls, data: dict[str, Any]) -> LicenseInfo:
66
+ """Deserialize from JSON cache."""
67
+ validated_at = data.get("validated_at")
68
+ if isinstance(validated_at, str):
69
+ data["validated_at"] = datetime.fromisoformat(validated_at)
70
+ trial_expires_at = data.get("trial_expires_at")
71
+ if isinstance(trial_expires_at, str):
72
+ data["trial_expires_at"] = datetime.fromisoformat(trial_expires_at)
73
+ elif trial_expires_at is None:
74
+ data["trial_expires_at"] = None
75
+ return cls(
76
+ status=data.get("status", "invalid"),
77
+ org_id=data.get("org_id"),
78
+ org_name=data.get("org_name"),
79
+ plan=data.get("plan"),
80
+ quota_tokens=data.get("quota_tokens"),
81
+ trial_expires_at=data.get("trial_expires_at"),
82
+ validated_at=data.get("validated_at", datetime.now(timezone.utc)),
83
+ )
84
+
85
+
86
+ class UsageReporter:
87
+ """Background license validator and aggregate usage reporter.
88
+
89
+ Validates the license key on startup, then periodically sends aggregate
90
+ usage stats to the Headroom cloud. If the cloud is unreachable, the proxy
91
+ continues to operate using cached license info (grace period: 7 days).
92
+
93
+ Never sends: message content, API keys, prompts, tool results, user data.
94
+ """
95
+
96
+ def __init__(
97
+ self,
98
+ license_key: str,
99
+ cloud_url: str = "https://app.headroomlabs.ai",
100
+ report_interval: int = 300,
101
+ cache_path: Path | None = None,
102
+ ):
103
+ self._license_key = license_key
104
+ self._cloud_url = cloud_url.rstrip("/")
105
+ self._report_interval = report_interval
106
+ self._cache_path = cache_path or LICENSE_CACHE_PATH
107
+ self._license_info: LicenseInfo | None = None
108
+ self._proxy: HeadroomProxy | None = None
109
+ self._task: asyncio.Task[None] | None = None
110
+ self._http_client: httpx.AsyncClient | None = None
111
+ self._stopped = False
112
+
113
+ # Snapshot of proxy metrics at last report (for computing deltas)
114
+ self._last_report_time: datetime | None = None
115
+ self._last_tokens_saved_by_model: dict[str, int] = {}
116
+ self._last_tokens_sent_by_model: dict[str, int] = {}
117
+ self._last_requests_by_model: dict[str, int] = {}
118
+
119
+ async def validate_license(self) -> LicenseInfo:
120
+ """Validate the license key against the cloud API.
121
+
122
+ On failure, falls back to cached license info if within grace period.
123
+ """
124
+ try:
125
+ client = await self._get_client()
126
+ resp = await client.post(
127
+ f"{self._cloud_url}/v1/license/validate",
128
+ json={"license_key": self._license_key},
129
+ timeout=10.0,
130
+ )
131
+ if resp.status_code == 200:
132
+ data = resp.json()
133
+ trial_expires_at = data.get("trial_expires_at")
134
+ if isinstance(trial_expires_at, str):
135
+ trial_expires_at = datetime.fromisoformat(trial_expires_at)
136
+
137
+ self._license_info = LicenseInfo(
138
+ status=data.get("status", "invalid"),
139
+ org_id=data.get("org_id"),
140
+ org_name=data.get("org_name"),
141
+ plan=data.get("plan"),
142
+ quota_tokens=data.get("quota_tokens"),
143
+ trial_expires_at=trial_expires_at,
144
+ validated_at=datetime.now(timezone.utc),
145
+ )
146
+ self._save_cache()
147
+ logger.info(
148
+ "License validated: status=%s org=%s plan=%s",
149
+ self._license_info.status,
150
+ self._license_info.org_name,
151
+ self._license_info.plan,
152
+ )
153
+ return self._license_info
154
+ else:
155
+ logger.warning(
156
+ "License validation returned status %d, using cached info",
157
+ resp.status_code,
158
+ )
159
+ except Exception:
160
+ logger.warning(
161
+ "Could not reach license server at %s, using cached info",
162
+ self._cloud_url,
163
+ exc_info=True,
164
+ )
165
+
166
+ # Fallback to cache
167
+ return self._load_cache_or_default()
168
+
169
+ async def start(self, proxy: HeadroomProxy) -> None:
170
+ """Start the background reporting loop. Called during proxy startup."""
171
+ self._proxy = proxy
172
+ self._stopped = False
173
+
174
+ # Validate license first
175
+ await self.validate_license()
176
+
177
+ # Take initial snapshot of metrics
178
+ self._snapshot_metrics()
179
+
180
+ # Start background reporting loop
181
+ self._task = asyncio.create_task(self._report_loop())
182
+ logger.info(
183
+ "Usage reporter started (interval=%ds, cloud=%s)",
184
+ self._report_interval,
185
+ self._cloud_url,
186
+ )
187
+
188
+ async def stop(self) -> None:
189
+ """Stop the background reporting loop. Called during proxy shutdown."""
190
+ self._stopped = True
191
+ if self._task and not self._task.done():
192
+ self._task.cancel()
193
+ try:
194
+ await self._task
195
+ except asyncio.CancelledError:
196
+ pass
197
+ if self._http_client:
198
+ await self._http_client.aclose()
199
+ self._http_client = None
200
+ logger.info("Usage reporter stopped")
201
+
202
+ @property
203
+ def is_active(self) -> bool:
204
+ """Whether the license is valid and usage is within quota."""
205
+ if self._license_info is None:
206
+ return True # No license info yet = allow (grace)
207
+ return self._license_info.status in ("active", "trial")
208
+
209
+ @property
210
+ def should_compress(self) -> bool:
211
+ """Whether compression should be applied. False = passthrough mode.
212
+
213
+ Returns True (allow compression) unless the license is definitively expired
214
+ AND outside the grace period.
215
+ """
216
+ if self._license_info is None:
217
+ return True # No license info yet = allow compression
218
+ if self._license_info.status in ("active", "trial"):
219
+ return True
220
+ if self._license_info.status == "expired":
221
+ # Check grace period
222
+ age = (datetime.now(timezone.utc) - self._license_info.validated_at).total_seconds()
223
+ if age < GRACE_PERIOD_SECONDS:
224
+ return True
225
+ return False
226
+ # "invalid" or unknown status: still allow compression (fail open)
227
+ return True
228
+
229
+ # ------------------------------------------------------------------
230
+ # Internal helpers
231
+ # ------------------------------------------------------------------
232
+
233
+ async def _get_client(self) -> httpx.AsyncClient:
234
+ if self._http_client is None:
235
+ self._http_client = httpx.AsyncClient(
236
+ timeout=httpx.Timeout(10.0),
237
+ headers={"User-Agent": "headroom-proxy"},
238
+ )
239
+ return self._http_client
240
+
241
+ async def _report_loop(self) -> None:
242
+ """Background loop: report usage every N seconds."""
243
+ while not self._stopped:
244
+ try:
245
+ await asyncio.sleep(self._report_interval)
246
+ if self._stopped:
247
+ break
248
+ await self._report_usage()
249
+ except asyncio.CancelledError:
250
+ break
251
+ except Exception:
252
+ logger.warning("Usage report failed, will retry next interval", exc_info=True)
253
+
254
+ async def _report_usage(self) -> None:
255
+ """Collect aggregate stats from the proxy and send to cloud."""
256
+ if self._proxy is None:
257
+ return
258
+
259
+ cost_tracker = self._proxy.cost_tracker
260
+ if cost_tracker is None:
261
+ return
262
+
263
+ now = datetime.now(timezone.utc)
264
+ period_start = self._last_report_time or now
265
+
266
+ # Compute deltas since last report
267
+ current_saved = dict(cost_tracker._tokens_saved_by_model)
268
+ current_sent = dict(cost_tracker._tokens_sent_by_model)
269
+ current_reqs = dict(cost_tracker._requests_by_model)
270
+
271
+ delta_saved_by_model: dict[str, int] = {}
272
+ delta_sent_by_model: dict[str, int] = {}
273
+ delta_reqs_by_model: dict[str, int] = {}
274
+
275
+ all_models = set(current_saved) | set(current_sent) | set(current_reqs)
276
+ total_tokens_saved = 0
277
+ total_tokens_before = 0
278
+ total_tokens_after = 0
279
+ total_requests = 0
280
+
281
+ for model in all_models:
282
+ saved = current_saved.get(model, 0) - self._last_tokens_saved_by_model.get(model, 0)
283
+ sent = current_sent.get(model, 0) - self._last_tokens_sent_by_model.get(model, 0)
284
+ reqs = current_reqs.get(model, 0) - self._last_requests_by_model.get(model, 0)
285
+ if reqs > 0:
286
+ delta_reqs_by_model[model] = reqs
287
+ if saved > 0:
288
+ delta_saved_by_model[model] = saved
289
+ if sent > 0:
290
+ delta_sent_by_model[model] = sent
291
+ total_tokens_saved += max(0, saved)
292
+ total_tokens_after += max(0, sent)
293
+ total_tokens_before += max(0, saved) + max(0, sent)
294
+ total_requests += max(0, reqs)
295
+
296
+ # Skip empty reports
297
+ if total_requests == 0:
298
+ self._last_report_time = now
299
+ return
300
+
301
+ payload = {
302
+ "license_key": self._license_key,
303
+ "period_start": period_start.isoformat(),
304
+ "period_end": now.isoformat(),
305
+ "requests": total_requests,
306
+ "tokens_before": total_tokens_before,
307
+ "tokens_after": total_tokens_after,
308
+ "tokens_saved": total_tokens_saved,
309
+ "models": delta_reqs_by_model,
310
+ }
311
+
312
+ try:
313
+ client = await self._get_client()
314
+ resp = await client.post(
315
+ f"{self._cloud_url}/v1/license/usage",
316
+ json=payload,
317
+ timeout=10.0,
318
+ )
319
+ if resp.status_code == 200:
320
+ data = resp.json()
321
+ status = data.get("status")
322
+ if status == "expired" and self._license_info:
323
+ self._license_info.status = "expired"
324
+ self._save_cache()
325
+ logger.warning("License expired: %s", data.get("message", ""))
326
+ elif status and self._license_info:
327
+ self._license_info.status = status
328
+ logger.debug(
329
+ "Usage reported: %d requests, %d tokens saved",
330
+ total_requests,
331
+ total_tokens_saved,
332
+ )
333
+ else:
334
+ logger.warning("Usage report returned status %d", resp.status_code)
335
+ except Exception:
336
+ logger.warning("Failed to send usage report", exc_info=True)
337
+
338
+ # Update snapshot
339
+ self._snapshot_metrics()
340
+ self._last_report_time = now
341
+
342
+ def _snapshot_metrics(self) -> None:
343
+ """Take a snapshot of current proxy metrics for delta computation."""
344
+ if self._proxy is None or self._proxy.cost_tracker is None:
345
+ return
346
+ ct = self._proxy.cost_tracker
347
+ self._last_tokens_saved_by_model = dict(ct._tokens_saved_by_model)
348
+ self._last_tokens_sent_by_model = dict(ct._tokens_sent_by_model)
349
+ self._last_requests_by_model = dict(ct._requests_by_model)
350
+ self._last_report_time = datetime.now(timezone.utc)
351
+
352
+ def _save_cache(self) -> None:
353
+ """Save license info to local cache file."""
354
+ if self._license_info is None:
355
+ return
356
+ try:
357
+ self._cache_path.parent.mkdir(parents=True, exist_ok=True)
358
+ self._cache_path.write_text(json.dumps(self._license_info.to_dict(), indent=2))
359
+ except OSError:
360
+ logger.warning("Could not save license cache to %s", self._cache_path)
361
+
362
+ def _load_cache_or_default(self) -> LicenseInfo:
363
+ """Load cached license info, or return a default if expired/missing."""
364
+ try:
365
+ if self._cache_path.exists():
366
+ data = json.loads(self._cache_path.read_text())
367
+ cached = LicenseInfo.from_dict(data)
368
+ age = (datetime.now(timezone.utc) - cached.validated_at).total_seconds()
369
+ if age < GRACE_PERIOD_SECONDS:
370
+ logger.info(
371
+ "Using cached license (age=%.1fh, status=%s)",
372
+ age / 3600,
373
+ cached.status,
374
+ )
375
+ self._license_info = cached
376
+ return cached
377
+ else:
378
+ logger.warning(
379
+ "Cached license expired (age=%.1fd), marking as expired",
380
+ age / 86400,
381
+ )
382
+ except (OSError, json.JSONDecodeError, KeyError):
383
+ logger.warning("Could not read license cache")
384
+
385
+ # No valid cache — return expired but still allow proxy to work
386
+ self._license_info = LicenseInfo(status="expired")
387
+ return self._license_info