speed up entitlement auth sync

This commit is contained in:
2569718930@qq.com
2026-06-10 16:49:18 +08:00
parent 7e17ab8a1b
commit 081bf9cc05
4 changed files with 470 additions and 42 deletions
+115 -21
View File
@@ -1430,28 +1430,44 @@ class SupabaseEntitlementService:
self, self,
user_id: str, user_id: str,
) -> Optional[Dict[str, object]]: ) -> Optional[Dict[str, object]]:
row, _ok = self._query_latest_active_subscription_result(user_id)
return row
def _query_latest_active_subscription_result(
self,
user_id: str,
bypass_cache: bool = False,
) -> Tuple[Optional[Dict[str, object]], bool]:
if not user_id: if not user_id:
return None return None, True
if not self.service_role_key: if not self.service_role_key:
logger.warning("SUPABASE_SERVICE_ROLE_KEY is missing") logger.warning("SUPABASE_SERVICE_ROLE_KEY is missing")
return None return None, False
now_ts = time.time() now_ts = time.time()
with self._sub_cache_lock: if not bypass_cache:
cached = self._sub_cache.get(user_id) with self._sub_cache_lock:
if cached and now_ts - float(cached.get("ts") or 0) < self.sub_cache_ttl_sec: cached = self._sub_cache.get(user_id)
row = cached.get("row") if cached and now_ts - float(cached.get("ts") or 0) < self.sub_cache_ttl_sec:
if isinstance(row, dict): rows = cached.get("rows")
return row if isinstance(rows, list):
return None return (
with self._active_subscription_bool_cache_lock: self._pick_latest_current_subscription(
cached_bool = self._active_subscription_bool_cache.get(user_id) [row for row in rows if isinstance(row, dict)]
if ( ),
cached_bool True,
and now_ts - float(cached_bool.get("ts") or 0) < self.sub_cache_ttl_sec )
and cached_bool.get("active") is False if "row" in cached:
): row = cached.get("row")
return None return row if isinstance(row, dict) else None, True
with self._active_subscription_bool_cache_lock:
cached_bool = self._active_subscription_bool_cache.get(user_id)
if (
cached_bool
and now_ts - float(cached_bool.get("ts") or 0) < self.sub_cache_ttl_sec
and cached_bool.get("active") is False
):
return None, True
try: try:
now = datetime.now(timezone.utc) now = datetime.now(timezone.utc)
@@ -1477,8 +1493,7 @@ class SupabaseEntitlementService:
user_id, user_id,
response.status_code, response.status_code,
) )
row = None return None, False
rows: List[Dict[str, object]] = []
else: else:
data = response.json() if response.content else [] data = response.json() if response.content else []
rows = [item for item in data if isinstance(item, dict)] if isinstance(data, list) else [] rows = [item for item in data if isinstance(item, dict)] if isinstance(data, list) else []
@@ -1490,10 +1505,89 @@ class SupabaseEntitlementService:
"row": row, "row": row,
"ts": now_ts, "ts": now_ts,
} }
return row return row, True
except Exception as exc: except Exception as exc:
logger.warning(f"supabase subscription query error user_id={user_id}: {exc}") logger.warning(f"supabase subscription query error user_id={user_id}: {exc}")
return None return None, False
def _subscription_access_window_from_current(
self,
row: Optional[Dict[str, object]],
) -> Dict[str, object]:
if not isinstance(row, dict):
return {}
expires_at = row.get("expires_at")
starts_at = row.get("starts_at")
return {
"current": row,
"current_expires_at": expires_at,
"current_starts_at": starts_at,
"total_expires_at": expires_at,
"queued_days": 0,
"queued_count": 0,
"rows": [row],
}
def _cached_subscription_access_window(
self,
user_id: str,
) -> Tuple[Dict[str, object], bool]:
now_ts = time.time()
with self._sub_cache_lock:
cached = self._sub_cache.get(user_id)
if not cached or now_ts - float(cached.get("ts") or 0) >= self.sub_cache_ttl_sec:
return {}, False
rows = cached.get("rows")
if isinstance(rows, list):
return (
self._subscription_window_from_rows(
[row for row in rows if isinstance(row, dict)]
),
True,
)
if "row" in cached:
row = cached.get("row")
return (
self._subscription_access_window_from_current(
row if isinstance(row, dict) else None
),
True,
)
return {}, False
def get_subscription_access_window(
self,
user_id: str,
respect_requirement: bool = True,
bypass_cache: bool = False,
unknown_on_error: bool = False,
) -> Dict[str, object]:
if respect_requirement and not self.require_subscription:
return {}
user_key = str(user_id or "").strip()
if not user_key:
return {}
if not bypass_cache:
cached_window, found = self._cached_subscription_access_window(user_key)
if found:
return cached_window
row, query_ok = self._query_latest_active_subscription_result(
user_key,
bypass_cache=True,
)
if not query_ok and unknown_on_error:
return {
"unknown": True,
"current": None,
"current_expires_at": None,
"current_starts_at": None,
"total_expires_at": None,
"queued_days": 0,
"queued_count": 0,
"rows": None,
}
return self._subscription_access_window_from_current(row)
def _query_active_subscription_rows_result( def _query_active_subscription_rows_result(
self, self,
+73
View File
@@ -1,4 +1,5 @@
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
import time
import src.auth.supabase_entitlement as entitlement_module import src.auth.supabase_entitlement as entitlement_module
from src.auth.supabase_entitlement import SupabaseEntitlementService from src.auth.supabase_entitlement import SupabaseEntitlementService
@@ -187,6 +188,78 @@ def test_subscription_window_can_report_unknown_on_transient_query_failure(monke
assert window["rows"] is None assert window["rows"] is None
def test_subscription_access_window_uses_single_current_subscription_query(monkeypatch):
monkeypatch.setenv("SUPABASE_URL", "https://example.supabase.co")
monkeypatch.setenv("SUPABASE_ANON_KEY", "anon-key")
monkeypatch.setenv("SUPABASE_SERVICE_ROLE_KEY", "service-role")
service = SupabaseEntitlementService()
def _fake_get(url, headers=None, params=None, timeout=None):
assert params["select"] == "plan_code,source,starts_at,expires_at"
assert str(params["starts_at"]).startswith("lte.")
assert params["limit"] == "1"
return _Response(
200,
[
{
"plan_code": "pro_monthly",
"source": "payment_contract",
"starts_at": "2026-03-01T00:00:00+00:00",
"expires_at": "2099-04-01T00:00:00+00:00",
}
],
)
monkeypatch.setattr(entitlement_module.requests, "get", _fake_get)
window = service.get_subscription_access_window(
"user-1",
respect_requirement=False,
unknown_on_error=True,
)
assert window["current"]["plan_code"] == "pro_monthly"
assert window["queued_days"] == 0
assert window["rows"] == [window["current"]]
def test_subscription_access_window_reuses_cached_current_row(monkeypatch):
monkeypatch.setenv("SUPABASE_URL", "https://example.supabase.co")
monkeypatch.setenv("SUPABASE_ANON_KEY", "anon-key")
monkeypatch.setenv("SUPABASE_SERVICE_ROLE_KEY", "service-role")
service = SupabaseEntitlementService()
current = {
"plan_code": "pro_monthly",
"source": "payment_contract",
"starts_at": "2026-03-01T00:00:00+00:00",
"expires_at": "2099-04-01T00:00:00+00:00",
}
service._sub_cache["user-1"] = {
"active": True,
"row": current,
"ts": time.time(),
}
monkeypatch.setattr(
entitlement_module.requests,
"get",
lambda *args, **kwargs: (_ for _ in ()).throw(
AssertionError("cached access window must not hit Supabase"),
),
)
window = service.get_subscription_access_window(
"user-1",
respect_requirement=False,
unknown_on_error=True,
)
assert window["current"] == current
assert window["total_expires_at"] == current["expires_at"]
def test_list_subscription_windows_selects_only_batch_window_fields(monkeypatch): def test_list_subscription_windows_selects_only_batch_window_fields(monkeypatch):
monkeypatch.setenv("SUPABASE_URL", "https://example.supabase.co") monkeypatch.setenv("SUPABASE_URL", "https://example.supabase.co")
monkeypatch.setenv("SUPABASE_ANON_KEY", "anon-key") monkeypatch.setenv("SUPABASE_ANON_KEY", "anon-key")
+168 -7
View File
@@ -2129,7 +2129,7 @@ def test_auth_me_uses_subscription_window_as_required_subscription_gate(monkeypa
assert payload["subscription_queued_days"] == 30 assert payload["subscription_queued_days"] == 30
def test_auth_me_entitlement_scope_reuses_subscription_window_cache(monkeypatch): def test_auth_me_entitlement_scope_reuses_subscription_access_window_cache(monkeypatch):
monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "enabled", True) monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "enabled", True)
monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "require_subscription", False) monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "require_subscription", False)
monkeypatch.setattr(web_core, "_SUPABASE_AUTH_REQUIRED", False) monkeypatch.setattr(web_core, "_SUPABASE_AUTH_REQUIRED", False)
@@ -2142,16 +2142,15 @@ def test_auth_me_entitlement_scope_reuses_subscription_window_cache(monkeypatch)
calls = [] calls = []
def _subscription_window( def _subscription_access_window(
user_id, user_id,
respect_requirement=False, respect_requirement=False,
bypass_cache=False,
unknown_on_error=False, unknown_on_error=False,
): ):
calls.append( calls.append(
{ {
"user_id": user_id, "user_id": user_id,
"bypass_cache": bypass_cache, "respect_requirement": respect_requirement,
"unknown_on_error": unknown_on_error, "unknown_on_error": unknown_on_error,
} }
) )
@@ -2166,7 +2165,12 @@ def test_auth_me_entitlement_scope_reuses_subscription_window_cache(monkeypatch)
monkeypatch.setattr(routes, "_assert_entitlement", lambda request: None) monkeypatch.setattr(routes, "_assert_entitlement", lambda request: None)
monkeypatch.setattr(routes, "_bind_optional_supabase_identity", _bind_identity) monkeypatch.setattr(routes, "_bind_optional_supabase_identity", _bind_identity)
monkeypatch.setattr(routes.SUPABASE_ENTITLEMENT, "get_subscription_window", _subscription_window) monkeypatch.setattr(
routes.SUPABASE_ENTITLEMENT,
"get_subscription_access_window",
_subscription_access_window,
raising=False,
)
response = client.get("/api/auth/me?scope=entitlement") response = client.get("/api/auth/me?scope=entitlement")
@@ -2174,7 +2178,163 @@ def test_auth_me_entitlement_scope_reuses_subscription_window_cache(monkeypatch)
assert calls == [ assert calls == [
{ {
"user_id": "user-1", "user_id": "user-1",
"bypass_cache": False, "respect_requirement": False,
"unknown_on_error": True,
}
]
def test_auth_me_entitlement_scope_defers_signup_trial_grant(monkeypatch):
monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "enabled", True)
monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "require_subscription", False)
monkeypatch.setattr(web_core, "_SUPABASE_AUTH_REQUIRED", False)
monkeypatch.setattr(routes, "_resolve_auth_points", lambda request: 0)
def _bind_identity(request):
request.state.auth_user_id = "user-1"
request.state.auth_email = "user@example.com"
request.state.auth_points = 0
scheduled = []
def _start_signup_trial_background(user_id, email):
scheduled.append((user_id, email))
return True
monkeypatch.setattr(routes, "_assert_entitlement", lambda request: None)
monkeypatch.setattr(routes, "_bind_optional_supabase_identity", _bind_identity)
monkeypatch.setattr(
auth_api,
"_start_signup_trial_background",
_start_signup_trial_background,
raising=False,
)
monkeypatch.setattr(
routes.SUPABASE_ENTITLEMENT,
"ensure_signup_trial",
lambda *args, **kwargs: (_ for _ in ()).throw(
AssertionError("entitlement scope must not block on signup trial writes"),
),
)
monkeypatch.setattr(
routes.SUPABASE_ENTITLEMENT,
"get_subscription_access_window",
lambda user_id, respect_requirement=False, bypass_cache=False, unknown_on_error=False: {},
raising=False,
)
response = client.get("/api/auth/me?scope=entitlement")
assert response.status_code == 200
payload = response.json()
assert payload["subscription_active"] is None
assert scheduled == [("user-1", "user@example.com")]
def test_signup_trial_background_grant_is_singleflight_and_cooled_down(monkeypatch):
with auth_api._SIGNUP_TRIAL_INFLIGHT_LOCK:
auth_api._SIGNUP_TRIAL_INFLIGHT.clear()
auth_api._SIGNUP_TRIAL_RECENT_ATTEMPTS.clear()
submitted = []
ensure_calls = []
class _FakeExecutor:
def submit(self, fn):
submitted.append(fn)
return object()
monkeypatch.setattr(auth_api, "_SIGNUP_TRIAL_EXECUTOR", _FakeExecutor())
monkeypatch.setattr(auth_api, "_signup_trial_background_cooldown_sec", lambda: 300.0)
monkeypatch.setattr(
routes.SUPABASE_ENTITLEMENT,
"ensure_signup_trial",
lambda user_id, email: ensure_calls.append((user_id, email)),
)
assert auth_api._start_signup_trial_background("user-1", "user@example.com") is True
assert auth_api._start_signup_trial_background("user-1", "user@example.com") is False
assert len(submitted) == 1
submitted[0]()
assert ensure_calls == [("user-1", "user@example.com")]
assert auth_api._start_signup_trial_background("user-1", "user@example.com") is False
assert len(submitted) == 1
def test_auth_me_entitlement_scope_uses_access_window_fast_path(monkeypatch):
monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "enabled", True)
monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "require_subscription", False)
monkeypatch.setattr(web_core, "_SUPABASE_AUTH_REQUIRED", False)
monkeypatch.setattr(routes, "_resolve_auth_points", lambda request: 0)
def _bind_identity(request):
request.state.auth_user_id = "user-1"
request.state.auth_email = "user@example.com"
request.state.auth_points = 0
fast_calls = []
def _access_window(user_id, respect_requirement=False, unknown_on_error=False):
fast_calls.append(
{
"user_id": user_id,
"respect_requirement": respect_requirement,
"unknown_on_error": unknown_on_error,
}
)
return {
"current": {
"plan_code": "pro_monthly",
"source": "payment_contract",
"starts_at": "2026-06-01T00:00:00+00:00",
"expires_at": "2026-07-01T00:00:00+00:00",
},
"total_expires_at": "2026-07-01T00:00:00+00:00",
"queued_days": 0,
"queued_count": 0,
}
monkeypatch.setattr(routes, "_assert_entitlement", lambda request: None)
monkeypatch.setattr(routes, "_bind_optional_supabase_identity", _bind_identity)
monkeypatch.setattr(
auth_api,
"_start_signup_trial_background",
lambda *args: (_ for _ in ()).throw(
AssertionError("active entitlement must not schedule signup trial work"),
),
raising=False,
)
monkeypatch.setattr(
routes.SUPABASE_ENTITLEMENT,
"ensure_signup_trial",
lambda *args, **kwargs: None,
)
monkeypatch.setattr(
routes.SUPABASE_ENTITLEMENT,
"get_subscription_access_window",
_access_window,
raising=False,
)
monkeypatch.setattr(
routes.SUPABASE_ENTITLEMENT,
"get_subscription_window",
lambda *args, **kwargs: (_ for _ in ()).throw(
AssertionError("entitlement scope should use the current-access fast path"),
),
)
response = client.get("/api/auth/me?scope=entitlement")
assert response.status_code == 200
payload = response.json()
assert payload["subscription_active"] is True
assert payload["subscription_plan_code"] == "pro_monthly"
assert fast_calls == [
{
"user_id": "user-1",
"respect_requirement": False,
"unknown_on_error": True, "unknown_on_error": True,
} }
] ]
@@ -2397,7 +2557,7 @@ def test_auth_me_entitlement_scope_skips_non_access_profile_sections(monkeypatch
monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "get_identity", lambda token: _Identity()) monkeypatch.setattr(web_core.SUPABASE_ENTITLEMENT, "get_identity", lambda token: _Identity())
monkeypatch.setattr( monkeypatch.setattr(
web_core.SUPABASE_ENTITLEMENT, web_core.SUPABASE_ENTITLEMENT,
"get_subscription_window", "get_subscription_access_window",
lambda user_id, respect_requirement=False, bypass_cache=False, unknown_on_error=False: { lambda user_id, respect_requirement=False, bypass_cache=False, unknown_on_error=False: {
"current": { "current": {
"plan_code": "pro_monthly", "plan_code": "pro_monthly",
@@ -2408,6 +2568,7 @@ def test_auth_me_entitlement_scope_skips_non_access_profile_sections(monkeypatch
"queued_days": 0, "queued_days": 0,
"queued_count": 0, "queued_count": 0,
}, },
raising=False,
) )
monkeypatch.setattr( monkeypatch.setattr(
web_core.SUPABASE_ENTITLEMENT, web_core.SUPABASE_ENTITLEMENT,
+114 -14
View File
@@ -2,7 +2,10 @@
from __future__ import annotations from __future__ import annotations
import os
import threading
import time import time
from concurrent.futures import ThreadPoolExecutor
from typing import Any, Callable, Dict, Optional, TypeVar from typing import Any, Callable, Dict, Optional, TypeVar
from fastapi import HTTPException, Request from fastapi import HTTPException, Request
@@ -17,6 +20,38 @@ import web.routes as legacy_routes
T = TypeVar("T") T = TypeVar("T")
def _signup_trial_background_workers() -> int:
try:
return max(
1,
int(os.getenv("POLYWEATHER_SIGNUP_TRIAL_BACKGROUND_WORKERS", "2") or "2"),
)
except Exception:
return 2
def _signup_trial_background_cooldown_sec() -> float:
try:
return max(
0.0,
float(
os.getenv("POLYWEATHER_SIGNUP_TRIAL_BACKGROUND_COOLDOWN_SEC", "300")
or "300"
),
)
except Exception:
return 300.0
_SIGNUP_TRIAL_EXECUTOR = ThreadPoolExecutor(
max_workers=_signup_trial_background_workers(),
thread_name_prefix="signup-trial",
)
_SIGNUP_TRIAL_INFLIGHT: set[str] = set()
_SIGNUP_TRIAL_RECENT_ATTEMPTS: Dict[str, float] = {}
_SIGNUP_TRIAL_INFLIGHT_LOCK = threading.Lock()
class _AuthMeTimer: class _AuthMeTimer:
def __init__(self, request: Request): def __init__(self, request: Request):
self.request = request self.request = request
@@ -99,6 +134,49 @@ def _is_entitlement_scope(request: Request) -> bool:
) )
def _start_signup_trial_background(user_id: str, email: Optional[str]) -> bool:
user_key = str(user_id or "").strip()
if not user_key:
return False
now_ts = time.time()
with _SIGNUP_TRIAL_INFLIGHT_LOCK:
if user_key in _SIGNUP_TRIAL_INFLIGHT:
return False
last_attempt = _SIGNUP_TRIAL_RECENT_ATTEMPTS.get(user_key)
if (
last_attempt is not None
and now_ts - last_attempt < _signup_trial_background_cooldown_sec()
):
return False
_SIGNUP_TRIAL_INFLIGHT.add(user_key)
_SIGNUP_TRIAL_RECENT_ATTEMPTS[user_key] = now_ts
if len(_SIGNUP_TRIAL_RECENT_ATTEMPTS) > 4096:
oldest_key = min(
_SIGNUP_TRIAL_RECENT_ATTEMPTS,
key=lambda item: _SIGNUP_TRIAL_RECENT_ATTEMPTS[item],
)
_SIGNUP_TRIAL_RECENT_ATTEMPTS.pop(oldest_key, None)
def _run() -> None:
try:
legacy_routes.SUPABASE_ENTITLEMENT.ensure_signup_trial(user_key, email)
except Exception as exc: # pragma: no cover - defensive background guard
logger.warning("signup trial background grant failed user_id={}: {}", user_key, exc)
finally:
with _SIGNUP_TRIAL_INFLIGHT_LOCK:
_SIGNUP_TRIAL_INFLIGHT.discard(user_key)
try:
_SIGNUP_TRIAL_EXECUTOR.submit(_run)
except Exception as exc:
with _SIGNUP_TRIAL_INFLIGHT_LOCK:
_SIGNUP_TRIAL_INFLIGHT.discard(user_key)
_SIGNUP_TRIAL_RECENT_ATTEMPTS.pop(user_key, None)
logger.warning("signup trial background submit failed user_id={}: {}", user_key, exc)
return False
return True
def _state_points(request: Request) -> int: def _state_points(request: Request) -> int:
try: try:
return max(0, int(getattr(request.state, "auth_points", 0) or 0)) return max(0, int(getattr(request.state, "auth_points", 0) or 0))
@@ -143,23 +221,37 @@ def get_auth_me_payload(request: Request) -> Dict[str, Any]:
if legacy_routes.SUPABASE_ENTITLEMENT.enabled and user_id: if legacy_routes.SUPABASE_ENTITLEMENT.enabled and user_id:
try: try:
timer.measure( if not entitlement_scope:
"ensure_signup_trial", timer.measure(
lambda: legacy_routes.SUPABASE_ENTITLEMENT.ensure_signup_trial( "ensure_signup_trial",
user_id, lambda: legacy_routes.SUPABASE_ENTITLEMENT.ensure_signup_trial(
email,
),
)
try:
subscription_window = timer.measure(
"subscription_window",
lambda: legacy_routes.SUPABASE_ENTITLEMENT.get_subscription_window(
user_id, user_id,
respect_requirement=False, email,
bypass_cache=False,
unknown_on_error=True,
), ),
) )
try:
if entitlement_scope and hasattr(
legacy_routes.SUPABASE_ENTITLEMENT,
"get_subscription_access_window",
):
subscription_window = timer.measure(
"subscription_window",
lambda: legacy_routes.SUPABASE_ENTITLEMENT.get_subscription_access_window(
user_id,
respect_requirement=False,
unknown_on_error=True,
),
)
else:
subscription_window = timer.measure(
"subscription_window",
lambda: legacy_routes.SUPABASE_ENTITLEMENT.get_subscription_window(
user_id,
respect_requirement=False,
bypass_cache=False,
unknown_on_error=True,
),
)
except TypeError: except TypeError:
subscription_window = timer.measure( subscription_window = timer.measure(
"subscription_window", "subscription_window",
@@ -207,6 +299,14 @@ def get_auth_me_payload(request: Request) -> Dict[str, Any]:
None if subscription_window_unknown else bool(latest_subscription) None if subscription_window_unknown else bool(latest_subscription)
) )
subscription_active_for_log = subscription_active subscription_active_for_log = subscription_active
if entitlement_scope and subscription_active is not True:
trial_grant_started = timer.measure(
"schedule_signup_trial",
lambda: _start_signup_trial_background(user_id, email),
)
if trial_grant_started and subscription_active is False:
subscription_active = None
subscription_active_for_log = None
if subscription_required and subscription_active is False: if subscription_required and subscription_active is False:
raise HTTPException(status_code=403, detail="Subscription required") raise HTTPException(status_code=403, detail="Subscription required")