diff --git a/src/auth/supabase_entitlement.py b/src/auth/supabase_entitlement.py index 1ab582ec..0de4621a 100644 --- a/src/auth/supabase_entitlement.py +++ b/src/auth/supabase_entitlement.py @@ -2,6 +2,7 @@ from __future__ import annotations import os import secrets +import hashlib import threading import time from dataclasses import dataclass @@ -240,6 +241,108 @@ class SupabaseEntitlementService: out.append(address) return out + @staticmethod + def _event_payload(row: Dict[str, object]) -> Dict[str, object]: + payload = row.get("payload") if isinstance(row, dict) else None + return dict(payload) if isinstance(payload, dict) else {} + + def _fetch_entitlement_events( + self, + *, + user_id: Optional[str] = None, + action: Optional[str] = None, + since: Optional[datetime] = None, + limit: int = 1000, + ) -> List[Dict[str, object]]: + params: Dict[str, Any] = { + "select": "id,user_id,action,payload,created_at", + "order": "created_at.desc", + "limit": str(max(1, min(int(limit or 1000), 5000))), + } + if user_id: + params["user_id"] = f"eq.{user_id}" + if action: + params["action"] = action + if since is not None: + params["created_at"] = f"gte.{self._to_iso(since)}" + rows = self._rest( + "GET", + "entitlement_events", + params=params, + allowed_status=[200], + ) + return [row for row in rows if isinstance(row, dict)] if isinstance(rows, list) else [] + + def _trial_claim_exists_in_events( + self, + *, + user_id: str, + email: str, + telegram_user_id: Optional[int], + wallet_addresses: List[str], + ) -> bool: + try: + rows = self._fetch_entitlement_events( + action="eq.signup_trial_claimed", + limit=2000, + ) + except Exception: + return False + wallet_set = { + str(address or "").strip().lower() + for address in wallet_addresses + if str(address or "").strip() + } + for row in rows: + payload = self._event_payload(row) + if str(row.get("user_id") or payload.get("user_id") or "").strip() == user_id: + return True + if email and str(payload.get("email") or "").strip().lower() == email: + return True + if telegram_user_id and str(payload.get("telegram_user_id") or "") == str(telegram_user_id): + return True + event_wallets = payload.get("wallet_addresses") + if isinstance(event_wallets, list): + event_wallet_set = { + str(address or "").strip().lower() + for address in event_wallets + if str(address or "").strip() + } + if wallet_set and wallet_set.intersection(event_wallet_set): + return True + return False + + def _record_signup_trial_claim_event( + self, + *, + user_id: str, + email: str, + telegram_user_id: Optional[int], + wallet_addresses: List[str], + claimed_at: datetime, + ) -> None: + self._rest( + "POST", + "entitlement_events", + payload={ + "user_id": user_id, + "action": "signup_trial_claimed", + "reason": "trial_dedupe", + "actor": "supabase_auth", + "payload": { + "user_id": user_id, + "email": email, + "telegram_user_id": telegram_user_id, + "wallet_addresses": wallet_addresses, + "claimed_at": self._to_iso(claimed_at), + "storage": "entitlement_events", + }, + "created_at": self._to_iso(claimed_at), + }, + prefer="return=minimal", + allowed_status=[201], + ) + def _trial_claim_exists( self, *, @@ -267,7 +370,12 @@ class SupabaseEntitlementService: if isinstance(rows, list) and rows: return True except Exception: - raise + return self._trial_claim_exists_in_events( + user_id=user_id, + email=email, + telegram_user_id=telegram_user_id, + wallet_addresses=wallet_addresses, + ) if not wallet_addresses: return False @@ -317,31 +425,49 @@ class SupabaseEntitlementService: "claimed_at": self._to_iso(now), "metadata": {"wallet_addresses": wallet_addresses}, } - claim_rows = self._rest( - "POST", - "trial_claims", - payload=claim_payload, - prefer="return=representation", - allowed_status=[200, 201], - ) claim_id = None - if isinstance(claim_rows, list) and claim_rows and isinstance(claim_rows[0], dict): - claim_id = claim_rows[0].get("id") - if wallet_addresses and claim_id is not None: - self._rest( + try: + claim_rows = self._rest( "POST", - "trial_claim_wallets", - payload=[ - { - "trial_claim_id": claim_id, - "wallet_address": address, - "created_at": self._to_iso(now), - } - for address in wallet_addresses - ], - prefer="return=minimal", - allowed_status=[201], + "trial_claims", + payload=claim_payload, + prefer="return=representation", + allowed_status=[200, 201], ) + if isinstance(claim_rows, list) and claim_rows and isinstance(claim_rows[0], dict): + claim_id = claim_rows[0].get("id") + except Exception: + self._record_signup_trial_claim_event( + user_id=user_key, + email=normalized_email, + telegram_user_id=telegram_user_id, + wallet_addresses=wallet_addresses, + claimed_at=now, + ) + if wallet_addresses and claim_id is not None: + try: + self._rest( + "POST", + "trial_claim_wallets", + payload=[ + { + "trial_claim_id": claim_id, + "wallet_address": address, + "created_at": self._to_iso(now), + } + for address in wallet_addresses + ], + prefer="return=minimal", + allowed_status=[201], + ) + except Exception: + self._record_signup_trial_claim_event( + user_id=user_key, + email=normalized_email, + telegram_user_id=telegram_user_id, + wallet_addresses=wallet_addresses, + claimed_at=now, + ) subscription_payload = { "user_id": user_key, @@ -412,6 +538,60 @@ class SupabaseEntitlementService: def _normalize_referral_code(value: Optional[str]) -> str: return "".join(str(value or "").strip().upper().split()) + @staticmethod + def _fallback_referral_code_for(user_id: str) -> str: + digest = hashlib.sha256( + f"polyweather-referral-v1:{user_id}".encode("utf-8") + ).hexdigest() + return f"PW{digest[:8].upper()}" + + def _find_referrer_by_fallback_code(self, code: str) -> Optional[str]: + normalized_code = self._normalize_referral_code(code) + try: + rows = self._rest( + "GET", + "profiles", + params={ + "select": "id", + "order": "created_at.asc", + "limit": "5000", + }, + allowed_status=[200], + ) + except Exception: + return None + if not isinstance(rows, list): + return None + for row in rows: + if not isinstance(row, dict): + continue + candidate = str(row.get("id") or "").strip() + if candidate and self._fallback_referral_code_for(candidate) == normalized_code: + return candidate + return None + + def _find_referrer_for_referral_code(self, code: str) -> Optional[str]: + normalized_code = self._normalize_referral_code(code) + try: + rows = self._rest( + "GET", + "referral_codes", + params={ + "select": "user_id,code,status", + "code": f"eq.{normalized_code}", + "status": "eq.active", + "limit": "1", + }, + allowed_status=[200], + ) + if isinstance(rows, list) and rows and isinstance(rows[0], dict): + referrer_user_id = str(rows[0].get("user_id") or "").strip() + if referrer_user_id: + return referrer_user_id + except Exception: + pass + return self._find_referrer_by_fallback_code(normalized_code) + def ensure_referral_code(self, user_id: str) -> Optional[Dict[str, object]]: user_key = str(user_id or "").strip() if not user_key or not self.service_role_key: @@ -454,8 +634,76 @@ class SupabaseEntitlementService: continue except Exception as exc: logger.warning("referral code ensure failed user_id={}: {}", user_key, exc) + return { + "code": self._fallback_referral_code_for(user_key), + "status": "active", + "storage": "derived", + } + + def _get_event_referral_attribution(self, user_id: str) -> Optional[Dict[str, object]]: + user_key = str(user_id or "").strip() + if not user_key: + return None + try: + rows = self._fetch_entitlement_events( + user_id=user_key, + action="in.(referral_attribution_created,referral_attribution_converted,referral_attribution_capped)", + limit=50, + ) + except Exception: + return None + for row in rows: + action = str(row.get("action") or "").strip() + payload = self._event_payload(row) + if str(payload.get("referred_user_id") or row.get("user_id") or "").strip() != user_key: + continue + if action in {"referral_attribution_converted", "referral_attribution_capped"}: + return None + if action == "referral_attribution_created": + return { + "id": payload.get("id") or row.get("id"), + "code": str(payload.get("code") or "").strip().upper(), + "referrer_user_id": str(payload.get("referrer_user_id") or "").strip(), + "referred_user_id": user_key, + "status": "pending", + "created_at": row.get("created_at"), + "_storage": "entitlement_events", + } return None + def _record_referral_attribution_event( + self, + *, + referrer_user_id: str, + referred_user_id: str, + code: str, + created_at: datetime, + ) -> Dict[str, object]: + attribution = { + "id": f"event:{referred_user_id}:{self._normalize_referral_code(code)}", + "code": self._normalize_referral_code(code), + "referrer_user_id": referrer_user_id, + "referred_user_id": referred_user_id, + "status": "pending", + "created_at": self._to_iso(created_at), + "_storage": "entitlement_events", + } + self._rest( + "POST", + "entitlement_events", + payload={ + "user_id": referred_user_id, + "action": "referral_attribution_created", + "reason": "invite_code", + "actor": "account_center", + "payload": attribution, + "created_at": self._to_iso(created_at), + }, + prefer="return=minimal", + allowed_status=[201], + ) + return attribution + def get_pending_referral_attribution(self, user_id: str) -> Optional[Dict[str, object]]: user_key = str(user_id or "").strip() if not user_key: @@ -474,10 +722,10 @@ class SupabaseEntitlementService: allowed_status=[200], ) except Exception: - return None + return self._get_event_referral_attribution(user_key) if isinstance(rows, list) and rows and isinstance(rows[0], dict): return rows[0] - return None + return self._get_event_referral_attribution(user_key) def _current_month_reward_rows(self, referrer_user_id: str) -> List[Dict[str, object]]: month_start = datetime.now(timezone.utc).replace( @@ -499,9 +747,39 @@ class SupabaseEntitlementService: }, allowed_status=[200], ) + except Exception: + return self._current_month_reward_event_rows(referrer_user_id, month_start) + table_rows = [row for row in rows if isinstance(row, dict)] if isinstance(rows, list) else [] + if table_rows: + return table_rows + return self._current_month_reward_event_rows(referrer_user_id, month_start) + + def _current_month_reward_event_rows( + self, + referrer_user_id: str, + month_start: datetime, + ) -> List[Dict[str, object]]: + try: + rows = self._fetch_entitlement_events( + user_id=referrer_user_id, + action="eq.referral_reward_granted", + since=month_start, + limit=100, + ) except Exception: return [] - return [row for row in rows if isinstance(row, dict)] if isinstance(rows, list) else [] + out: List[Dict[str, object]] = [] + for row in rows: + payload = self._event_payload(row) + out.append( + { + "id": row.get("id"), + "reward_days": int(payload.get("reward_days") or REFERRAL_REWARD_DAYS), + "created_at": row.get("created_at"), + "_storage": "entitlement_events", + } + ) + return out def get_referral_summary(self, user_id: str) -> Optional[Dict[str, object]]: user_key = str(user_id or "").strip() @@ -539,20 +817,9 @@ class SupabaseEntitlementService: if self.has_paid_subscription(user_key): raise ValueError("referral code can only be used before first paid subscription") - rows = self._rest( - "GET", - "referral_codes", - params={ - "select": "user_id,code,status", - "code": f"eq.{normalized_code}", - "status": "eq.active", - "limit": "1", - }, - allowed_status=[200], - ) - if not isinstance(rows, list) or not rows or not isinstance(rows[0], dict): + referrer_user_id = self._find_referrer_for_referral_code(normalized_code) + if not referrer_user_id: raise ValueError("referral code not found") - referrer_user_id = str(rows[0].get("user_id") or "").strip() if not referrer_user_id or referrer_user_id == user_key: raise ValueError("cannot use your own referral code") @@ -565,20 +832,28 @@ class SupabaseEntitlementService: } now = datetime.now(timezone.utc) - self._rest( - "POST", - "referral_attributions", - payload={ - "referrer_user_id": referrer_user_id, - "referred_user_id": user_key, - "code": normalized_code, - "status": "pending", - "created_at": self._to_iso(now), - "updated_at": self._to_iso(now), - }, - prefer="return=minimal", - allowed_status=[201], - ) + try: + self._rest( + "POST", + "referral_attributions", + payload={ + "referrer_user_id": referrer_user_id, + "referred_user_id": user_key, + "code": normalized_code, + "status": "pending", + "created_at": self._to_iso(now), + "updated_at": self._to_iso(now), + }, + prefer="return=minimal", + allowed_status=[201], + ) + except Exception: + self._record_referral_attribution_event( + referrer_user_id=referrer_user_id, + referred_user_id=user_key, + code=normalized_code, + created_at=now, + ) return { "ok": True, "already_applied": False, @@ -615,6 +890,113 @@ class SupabaseEntitlementService: break return starts + def _record_referral_resolution_event( + self, + *, + action: str, + referrer_user_id: str, + referred_user_id: str, + attribution: Dict[str, object], + payment_intent_id: str, + tx_hash: str, + created_at: datetime, + reward_days: int = 0, + ) -> None: + self._rest( + "POST", + "entitlement_events", + payload={ + "user_id": referrer_user_id if action == "referral_reward_granted" else referred_user_id, + "action": action, + "reason": "referred_user_paid", + "actor": "payment_contract_checkout", + "payload": { + "attribution_id": attribution.get("id"), + "code": attribution.get("code"), + "referrer_user_id": referrer_user_id, + "referred_user_id": referred_user_id, + "payment_intent_id": payment_intent_id, + "tx_hash": tx_hash, + "reward_days": reward_days, + "storage": "entitlement_events", + }, + "created_at": self._to_iso(created_at), + }, + prefer="return=minimal", + allowed_status=[201], + ) + + def _settle_referral_reward_with_events( + self, + *, + attribution: Dict[str, object], + referrer_user_id: str, + referred_user_id: str, + payment_intent_id: str, + tx_hash: str, + now: datetime, + monthly_rewards: List[Dict[str, object]], + ) -> Dict[str, object]: + if len(monthly_rewards) >= REFERRAL_MONTHLY_REWARD_LIMIT: + self._record_referral_resolution_event( + action="referral_attribution_capped", + referrer_user_id=referrer_user_id, + referred_user_id=referred_user_id, + attribution=attribution, + payment_intent_id=payment_intent_id, + tx_hash=tx_hash, + created_at=now, + ) + return {"awarded": False, "reason": "monthly_cap_reached"} + + starts = self._subscription_extension_start(referrer_user_id) + expires = starts + timedelta(days=REFERRAL_REWARD_DAYS) + subscription_payload = { + "user_id": referrer_user_id, + "plan_code": "pro_monthly", + "status": "active", + "starts_at": self._to_iso(starts), + "expires_at": self._to_iso(expires), + "source": "referral_reward", + "created_at": self._to_iso(now), + "updated_at": self._to_iso(now), + } + self._rest( + "POST", + "subscriptions", + payload=subscription_payload, + prefer="return=minimal", + allowed_status=[201], + ) + self._record_referral_resolution_event( + action="referral_reward_granted", + referrer_user_id=referrer_user_id, + referred_user_id=referred_user_id, + attribution=attribution, + payment_intent_id=payment_intent_id, + tx_hash=tx_hash, + created_at=now, + reward_days=REFERRAL_REWARD_DAYS, + ) + self._record_referral_resolution_event( + action="referral_attribution_converted", + referrer_user_id=referrer_user_id, + referred_user_id=referred_user_id, + attribution=attribution, + payment_intent_id=payment_intent_id, + tx_hash=tx_hash, + created_at=now, + reward_days=REFERRAL_REWARD_DAYS, + ) + self.invalidate_subscription_cache(referrer_user_id) + return { + "awarded": True, + "reward_days": REFERRAL_REWARD_DAYS, + "referrer_user_id": referrer_user_id, + "subscription": subscription_payload, + "storage": "entitlement_events", + } + def settle_referral_reward( self, *, @@ -632,6 +1014,16 @@ class SupabaseEntitlementService: now = datetime.now(timezone.utc) monthly_rewards = self._current_month_reward_rows(referrer_key) + if str(attribution.get("_storage") or "") == "entitlement_events": + return self._settle_referral_reward_with_events( + attribution=attribution, + referrer_user_id=referrer_key, + referred_user_id=referred_key, + payment_intent_id=payment_intent_id, + tx_hash=tx_hash, + now=now, + monthly_rewards=monthly_rewards, + ) if len(monthly_rewards) >= REFERRAL_MONTHLY_REWARD_LIMIT: self._rest( "PATCH", diff --git a/tests/test_subscription_referral_program.py b/tests/test_subscription_referral_program.py index 059c22e4..17b753fc 100644 --- a/tests/test_subscription_referral_program.py +++ b/tests/test_subscription_referral_program.py @@ -177,3 +177,133 @@ def test_referral_reward_respects_monthly_ten_invite_cap(monkeypatch): assert result["awarded"] is False assert result["reason"] == "monthly_cap_reached" assert not any(call["table"] == "subscriptions" for call in writes) + + +def test_signup_trial_falls_back_to_entitlement_events_without_trial_tables(monkeypatch): + monkeypatch.setenv("POLYWEATHER_AUTH_ENABLED", "true") + monkeypatch.setenv("SUPABASE_URL", "https://example.supabase.co") + monkeypatch.setenv("SUPABASE_ANON_KEY", "anon") + monkeypatch.setenv("SUPABASE_SERVICE_ROLE_KEY", "service-role") + service = SupabaseEntitlementService() + calls = [] + + def fake_rest(method, table, **kwargs): + calls.append({"method": method, "table": table, **kwargs}) + if method == "GET" and table == "user_wallets": + return [] + if table in {"trial_claims", "trial_claim_wallets"}: + raise RuntimeError("table missing") + if method == "GET" and table == "entitlement_events": + return [] + if method == "POST" and table in {"subscriptions", "entitlement_events"}: + return [kwargs.get("payload") or {}] + raise AssertionError((method, table, kwargs)) + + monkeypatch.setattr(service, "_rest", fake_rest) + + result = service.ensure_signup_trial("user-1", "user@example.com") + + assert result["created"] is True + event_actions = [ + call["payload"]["action"] + for call in calls + if call["method"] == "POST" and call["table"] == "entitlement_events" + ] + assert "signup_trial_claimed" in event_actions + assert "signup_trial_granted" in event_actions + + +def test_apply_referral_code_falls_back_to_profiles_and_events(monkeypatch): + monkeypatch.setenv("POLYWEATHER_AUTH_ENABLED", "true") + monkeypatch.setenv("SUPABASE_URL", "https://example.supabase.co") + monkeypatch.setenv("SUPABASE_ANON_KEY", "anon") + monkeypatch.setenv("SUPABASE_SERVICE_ROLE_KEY", "service-role") + service = SupabaseEntitlementService() + referrer_id = "00000000-0000-0000-0000-0000000000aa" + referred_id = "00000000-0000-0000-0000-0000000000bb" + events = [] + + def fake_rest(method, table, **kwargs): + if method == "GET" and table == "subscriptions": + return [] + if table in {"referral_codes", "referral_attributions", "referral_rewards"}: + raise RuntimeError("table missing") + if method == "GET" and table == "profiles": + return [{"id": referrer_id, "email": "referrer@example.com"}] + if method == "GET" and table == "entitlement_events": + user_filter = str((kwargs.get("params") or {}).get("user_id") or "") + user_id = user_filter[3:] if user_filter.startswith("eq.") else user_filter + return [event for event in events if event.get("user_id") == user_id] + if method == "POST" and table == "entitlement_events": + payload = kwargs.get("payload") or {} + events.append(payload) + return [payload] + raise AssertionError((method, table, kwargs)) + + monkeypatch.setattr(service, "_rest", fake_rest) + referral_code = service.ensure_referral_code(referrer_id)["code"] + + result = service.apply_referral_code(referred_id, referral_code) + + assert result["ok"] is True + assert result["referral"]["applied_code"] == referral_code + assert any(event["action"] == "referral_attribution_created" for event in events) + + +def test_referral_reward_falls_back_to_entitlement_events_without_referral_tables(monkeypatch): + monkeypatch.setenv("POLYWEATHER_AUTH_ENABLED", "true") + monkeypatch.setenv("SUPABASE_URL", "https://example.supabase.co") + monkeypatch.setenv("SUPABASE_ANON_KEY", "anon") + monkeypatch.setenv("SUPABASE_SERVICE_ROLE_KEY", "service-role") + service = SupabaseEntitlementService() + referrer_id = "00000000-0000-0000-0000-0000000000aa" + referred_id = "00000000-0000-0000-0000-0000000000bb" + events = [ + { + "id": 1, + "user_id": referred_id, + "action": "referral_attribution_created", + "payload": { + "id": "event-1", + "code": "PWTEST", + "referrer_user_id": referrer_id, + "referred_user_id": referred_id, + "status": "pending", + }, + "created_at": "2026-05-29T00:00:00+00:00", + } + ] + writes = [] + + def fake_rest(method, table, **kwargs): + if table in {"referral_attributions", "referral_rewards"}: + raise RuntimeError("table missing") + if method == "GET" and table == "entitlement_events": + user_filter = str((kwargs.get("params") or {}).get("user_id") or "") + user_id = user_filter[3:] if user_filter.startswith("eq.") else user_filter + return [event for event in events if event.get("user_id") == user_id] + if method == "GET" and table == "subscriptions": + return [] + if method == "POST" and table in {"subscriptions", "entitlement_events"}: + payload = kwargs.get("payload") or {} + writes.append({"table": table, "payload": payload}) + if table == "entitlement_events": + events.append(payload) + return [payload] + raise AssertionError((method, table, kwargs)) + + monkeypatch.setattr(service, "_rest", fake_rest) + + result = service.settle_referral_reward( + referred_user_id=referred_id, + payment_intent_id="intent-1", + tx_hash="0x" + "3" * 64, + ) + + assert result["awarded"] is True + assert any(call["table"] == "subscriptions" for call in writes) + assert any( + call["table"] == "entitlement_events" + and call["payload"]["action"] == "referral_reward_granted" + for call in writes + )