diff --git a/src/auth/supabase_entitlement.py b/src/auth/supabase_entitlement.py index a987d4de..68053718 100644 --- a/src/auth/supabase_entitlement.py +++ b/src/auth/supabase_entitlement.py @@ -4,7 +4,7 @@ import os import threading import time from dataclasses import dataclass -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from typing import Dict, List, Optional import requests @@ -42,6 +42,7 @@ class SupabaseIdentity: user_id: str email: str points: int = 0 + created_at: Optional[str] = None class SupabaseEntitlementService: @@ -64,11 +65,27 @@ class SupabaseEntitlementService: self.timeout_sec = max(3, _env_int("SUPABASE_HTTP_TIMEOUT_SEC", 8)) self.cache_ttl_sec = max(5, _env_int("SUPABASE_AUTH_CACHE_TTL_SEC", 30)) self.sub_cache_ttl_sec = max(5, _env_int("SUPABASE_SUB_CACHE_TTL_SEC", 60)) + self.signup_trial_enabled = _env_bool( + "POLYWEATHER_SIGNUP_TRIAL_ENABLED", + True, + ) + self.signup_trial_days = max( + 0, + _env_int("POLYWEATHER_SIGNUP_TRIAL_DAYS", 3), + ) + self.signup_trial_plan_code = str( + os.getenv("POLYWEATHER_SIGNUP_TRIAL_PLAN_CODE") or "signup_trial_3d" + ).strip() or "signup_trial_3d" + self.signup_trial_source = str( + os.getenv("POLYWEATHER_SIGNUP_TRIAL_SOURCE") or "signup_trial" + ).strip() or "signup_trial" self._identity_cache: Dict[str, Dict[str, object]] = {} self._identity_cache_lock = threading.Lock() self._sub_cache: Dict[str, Dict[str, object]] = {} self._sub_cache_lock = threading.Lock() + self._trial_locks: Dict[str, threading.Lock] = {} + self._trial_locks_guard = threading.Lock() def invalidate_subscription_cache(self, user_id: str) -> None: key = str(user_id or "").strip() @@ -87,6 +104,9 @@ class SupabaseEntitlementService: def _subscription_endpoint(self) -> str: return f"{self.supabase_url}/rest/v1/subscriptions" + def _entitlement_events_endpoint(self) -> str: + return f"{self.supabase_url}/rest/v1/entitlement_events" + def _request_headers_for_user(self, access_token: str) -> Dict[str, str]: return { "apikey": self.anon_key, @@ -140,6 +160,7 @@ class SupabaseEntitlementService: user_id=user_id, email=str(data.get("email") or "").strip(), points=points, + created_at=str(data.get("created_at") or "").strip() or None, ) with self._identity_cache_lock: self._identity_cache[access_token] = { @@ -213,6 +234,180 @@ class SupabaseEntitlementService: logger.warning(f"supabase subscription query error user_id={user_id}: {exc}") return None + def _query_latest_subscription_any_status( + self, + user_id: str, + ) -> Optional[Dict[str, object]]: + if not user_id or not self.service_role_key: + return None + try: + params = { + "select": "id,user_id,status,plan_code,starts_at,expires_at,source,created_at,updated_at", + "user_id": f"eq.{user_id}", + "order": "created_at.desc", + "limit": "1", + } + response = requests.get( + self._subscription_endpoint(), + headers=self._request_headers_for_service_role(), + params=params, + timeout=self.timeout_sec, + ) + if response.status_code != 200: + logger.warning( + "supabase subscription history query failed user_id={} status={}", + user_id, + response.status_code, + ) + return None + data = response.json() if response.content else [] + row = data[0] if isinstance(data, list) and data else None + return row if isinstance(row, dict) else None + except Exception as exc: + logger.warning(f"supabase subscription history query error user_id={user_id}: {exc}") + return None + + @staticmethod + def _parse_iso_datetime(raw: Optional[str]) -> Optional[datetime]: + text = str(raw or "").strip() + if not text: + return None + try: + parsed = datetime.fromisoformat(text.replace("Z", "+00:00")) + except Exception: + return None + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=timezone.utc) + return parsed.astimezone(timezone.utc) + + def _get_trial_lock(self, user_id: str) -> threading.Lock: + key = str(user_id or "").strip() + with self._trial_locks_guard: + lock = self._trial_locks.get(key) + if lock is None: + lock = threading.Lock() + self._trial_locks[key] = lock + return lock + + def _emit_signup_trial_event( + self, + *, + user_id: str, + starts_at: datetime, + expires_at: datetime, + ) -> None: + if not self.service_role_key: + return + try: + now_iso = datetime.now(timezone.utc).isoformat() + requests.post( + self._entitlement_events_endpoint(), + headers={ + **self._request_headers_for_service_role(), + "Content-Type": "application/json", + "Prefer": "return=minimal", + }, + json={ + "user_id": user_id, + "action": "subscription_granted", + "reason": "signup_trial", + "detail": f"{self.signup_trial_days}d signup trial granted", + "payload": { + "plan_code": self.signup_trial_plan_code, + "source": self.signup_trial_source, + "starts_at": starts_at.isoformat(), + "expires_at": expires_at.isoformat(), + }, + "created_at": now_iso, + }, + timeout=self.timeout_sec, + ) + except Exception as exc: + logger.warning(f"supabase signup trial event insert failed user_id={user_id}: {exc}") + + def ensure_signup_trial( + self, + user_id: str, + *, + created_at: Optional[str] = None, + ) -> Optional[Dict[str, object]]: + normalized_user_id = str(user_id or "").strip() + if not normalized_user_id: + return None + if ( + not self.signup_trial_enabled + or self.signup_trial_days <= 0 + or not self.service_role_key + ): + return None + + lock = self._get_trial_lock(normalized_user_id) + with lock: + existing_active = self._query_latest_active_subscription(normalized_user_id) + if isinstance(existing_active, dict): + return existing_active + + existing_any = self._query_latest_subscription_any_status(normalized_user_id) + if isinstance(existing_any, dict): + return None + + starts_at = self._parse_iso_datetime(created_at) + if starts_at is None: + admin_users = self.get_auth_users([normalized_user_id]) + starts_at = self._parse_iso_datetime( + str((admin_users.get(normalized_user_id) or {}).get("created_at") or "") + ) + if starts_at is None: + return None + + expires_at = starts_at + timedelta(days=self.signup_trial_days) + now = datetime.now(timezone.utc) + if expires_at <= now: + return None + + payload = { + "user_id": normalized_user_id, + "plan_code": self.signup_trial_plan_code, + "status": "active", + "starts_at": starts_at.isoformat(), + "expires_at": expires_at.isoformat(), + "source": self.signup_trial_source, + "created_at": now.isoformat(), + "updated_at": now.isoformat(), + } + try: + response = requests.post( + self._subscription_endpoint(), + headers={ + **self._request_headers_for_service_role(), + "Content-Type": "application/json", + "Prefer": "return=representation", + }, + json=payload, + timeout=self.timeout_sec, + ) + if response.status_code not in (200, 201): + logger.warning( + "supabase signup trial insert failed user_id={} status={}", + normalized_user_id, + response.status_code, + ) + return self._query_latest_active_subscription(normalized_user_id) + rows = response.json() if response.content else [] + row = rows[0] if isinstance(rows, list) and rows else None + self.invalidate_subscription_cache(normalized_user_id) + self._emit_signup_trial_event( + user_id=normalized_user_id, + starts_at=starts_at, + expires_at=expires_at, + ) + if isinstance(row, dict): + return row + return self._query_latest_active_subscription(normalized_user_id) + except Exception as exc: + logger.warning(f"supabase signup trial insert error user_id={normalized_user_id}: {exc}") + return self._query_latest_active_subscription(normalized_user_id) + def _query_active_subscription(self, user_id: str) -> bool: return self._query_latest_active_subscription(user_id) is not None diff --git a/tests/test_supabase_entitlement.py b/tests/test_supabase_entitlement.py new file mode 100644 index 00000000..bbaf5c87 --- /dev/null +++ b/tests/test_supabase_entitlement.py @@ -0,0 +1,66 @@ +from datetime import datetime, timedelta, timezone + +import src.auth.supabase_entitlement as entitlement_module +from src.auth.supabase_entitlement import SupabaseEntitlementService + + +class _Response: + def __init__(self, status_code=200, payload=None): + self.status_code = status_code + self._payload = payload + self.content = b"1" + + def json(self): + return self._payload + + +def test_ensure_signup_trial_grants_three_day_subscription(monkeypatch): + monkeypatch.setenv("SUPABASE_URL", "https://example.supabase.co") + monkeypatch.setenv("SUPABASE_ANON_KEY", "anon-key") + monkeypatch.setenv("SUPABASE_SERVICE_ROLE_KEY", "service-role") + monkeypatch.setenv("POLYWEATHER_SIGNUP_TRIAL_ENABLED", "true") + monkeypatch.setenv("POLYWEATHER_SIGNUP_TRIAL_DAYS", "3") + + service = SupabaseEntitlementService() + monkeypatch.setattr(service, "_query_latest_active_subscription", lambda user_id: None) + monkeypatch.setattr(service, "_query_latest_subscription_any_status", lambda user_id: None) + + captured_posts = [] + + def _fake_post(url, headers=None, json=None, timeout=None): + captured_posts.append({"url": url, "headers": headers, "json": json, "timeout": timeout}) + if url.endswith("/rest/v1/subscriptions"): + return _Response( + 201, + [ + { + "user_id": json["user_id"], + "plan_code": json["plan_code"], + "status": json["status"], + "starts_at": json["starts_at"], + "expires_at": json["expires_at"], + "source": json["source"], + } + ], + ) + return _Response(201, {}) + + monkeypatch.setattr(entitlement_module.requests, "post", _fake_post) + + starts_at = datetime(2026, 3, 29, 8, 0, tzinfo=timezone.utc) + result = service.ensure_signup_trial( + "user-1", + created_at=starts_at.isoformat(), + ) + + assert result is not None + assert result["plan_code"] == "signup_trial_3d" + assert result["status"] == "active" + assert result["starts_at"] == starts_at.isoformat() + assert result["expires_at"] == (starts_at + timedelta(days=3)).isoformat() + + subscription_insert = next( + item for item in captured_posts if item["url"].endswith("/rest/v1/subscriptions") + ) + assert subscription_insert["json"]["user_id"] == "user-1" + assert subscription_insert["json"]["source"] == "signup_trial" diff --git a/web/core.py b/web/core.py index adc9b20e..51074f3a 100644 --- a/web/core.py +++ b/web/core.py @@ -179,6 +179,7 @@ def _bind_optional_supabase_identity(request: Request) -> None: request.state.auth_user_id = identity.user_id request.state.auth_email = identity.email request.state.auth_points = identity.points + request.state.auth_created_at = identity.created_at def _resolve_auth_points(request: Request) -> int: diff --git a/web/routes.py b/web/routes.py index e3f28a7c..2e0da1c1 100644 --- a/web/routes.py +++ b/web/routes.py @@ -275,10 +275,15 @@ async def auth_me(request: Request): if SUPABASE_ENTITLEMENT.enabled and user_id: try: - latest_subscription = SUPABASE_ENTITLEMENT.get_latest_active_subscription( + latest_subscription = SUPABASE_ENTITLEMENT.ensure_signup_trial( user_id, - respect_requirement=False, + created_at=getattr(request.state, "auth_created_at", None), ) + if not latest_subscription: + latest_subscription = SUPABASE_ENTITLEMENT.get_latest_active_subscription( + user_id, + respect_requirement=False, + ) if ( not latest_subscription and getattr(PAYMENT_CHECKOUT, "enabled", False)