Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,9 @@ HCAPTCHA_SITEKEY="" # hCaptcha site key (public, used in temp
SECRET_KEY="" # To generate: python -c "import os; print(os.urandom(32).hex())"
HOST_URI="127.0.0.1:8000"
ENV="development" # change to "production" in production
# none = self-host (every account holds every feature); paddle = spoo.me cloud.
# Production refuses to boot unless this is set explicitly.
BILLING_PROVIDER="none"

# CORS — allowed origins for private routes (auth, oauth, dashboard)
# Public API routes (/api/v1/*) always allow all origins.
Expand Down
4 changes: 4 additions & 0 deletions app.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
from infrastructure.oauth_clients import OAUTH_STATE_TTL_SECONDS, init_oauth
from infrastructure.queue_redis import connect_queue_redis
from infrastructure.templates import configure_template_globals, templates
from middleware.entitlements import EntitlementsVersionMiddleware
from middleware.error_handler import register_error_handlers
from middleware.logging import RequestLoggingMiddleware
from middleware.openapi import (
Expand Down Expand Up @@ -314,6 +315,9 @@ async def docs(request: Request):
# writes view_rate_limit into shared scope state during endpoint
# execution, before any response starts flowing outward
app.add_middleware(RateLimitHeadersMiddleware)
# 8. Entitlement version header on authenticated responses; reads the
# version the Entitled dependency stored, or one cache lookup.
app.add_middleware(EntitlementsVersionMiddleware)

# ── Error handlers + rate limiter ────────────────────────────────────
app.state.limiter = limiter
Expand Down
35 changes: 35 additions & 0 deletions config.py
Original file line number Diff line number Diff line change
Expand Up @@ -664,6 +664,31 @@ class SchedulerSettings(BaseSettings):
lease_seconds: int = Field(default=600, ge=30)


class BillingSettings(BaseSettings):
"""Billing provider and the display prices the plans endpoint shows.

``BILLING_PROVIDER=none`` is self-host: no billing, and the resolver hands
every account the ``selfhost`` plan. Production must set it explicitly so
the cloud can never fall into self-host by omission. Prices are display
values only; the provider owns what is charged.
"""

model_config = SettingsConfigDict(
env_file=".env", env_prefix="BILLING_", extra="ignore"
)

provider: Literal["none", "paddle"] = "none"
pro_monthly_usd: int = Field(default=15, ge=0)
pro_year_usd: int = Field(default=144, ge=0)
founding_monthly_usd: int = Field(default=9, ge=0)
founding_year_usd: int = Field(default=90, ge=0)
founding_seats: int = Field(default=100, ge=0)

@property
def selfhost(self) -> bool:
return self.provider == "none"


class AppSettings(BaseSettings):
model_config = SettingsConfigDict(env_file=".env", extra="ignore")

Expand Down Expand Up @@ -859,6 +884,7 @@ def _password_max_length_sane(cls, v: int) -> int:
scheduler: SchedulerSettings | None = None
llm: LlmSettings | None = None
posthog_erasure: PostHogErasureSettings | None = None
billing: BillingSettings | None = None

@model_validator(mode="after")
def _populate_sub_configs_and_secret(self) -> AppSettings:
Expand Down Expand Up @@ -908,6 +934,8 @@ def _populate_sub_configs_and_secret(self) -> AppSettings:
self.scheduler = SchedulerSettings()
if self.posthog_erasure is None:
self.posthog_erasure = PostHogErasureSettings()
if self.billing is None:
self.billing = BillingSettings()
if self.webhooks.enabled and not self.secret_key:
# Signing secrets are encrypted with a key derived from
# SECRET_KEY; an empty master would mean a predictable key.
Expand All @@ -926,6 +954,13 @@ def _populate_sub_configs_and_secret(self) -> AppSettings:
if self.env == "production" and self.account_deletion_grace_days < 1:
raise ValueError("ACCOUNT_DELETION_GRACE_DAYS must be >= 1 in production")

# Unset means self-host, which hands every account every feature.
if self.env == "production" and "provider" not in self.billing.model_fields_set:
raise ValueError(
"BILLING_PROVIDER must be set explicitly in production: "
"none for self-host, paddle for the cloud"
)

return self

@property
Expand Down
10 changes: 10 additions & 0 deletions dependencies/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,12 @@
require_session_or_scopes,
require_verified_email,
)
from dependencies.entitlements import (
Entitled,
EntitlementSvc,
get_entitlement_service,
get_entitlements,
)
from dependencies.infra import (
AppRegistryDep,
GeoIP,
Expand Down Expand Up @@ -132,6 +138,8 @@
"CustomDomainSvc",
"DeviceAuthSvc",
"DomainIntelSvc",
"Entitled",
"EntitlementSvc",
"ExportSvc",
"FeatureFlagSvc",
"GeoIP",
Expand Down Expand Up @@ -175,6 +183,8 @@
"get_db",
"get_device_auth_service",
"get_email_provider",
"get_entitlement_service",
"get_entitlements",
"get_export_service",
"get_feature_flag_service",
"get_geoip_service",
Expand Down
12 changes: 4 additions & 8 deletions dependencies/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,10 +64,9 @@ class CurrentUser:
# access tokens minted before the claim existed; those users match by
# user_id only until their next token refresh.
email: str | None = field(default=None)
# UserDoc.plan value (e.g. "FREE") — consumed by FeatureFlagService's
# TIER rollout via getattr(user, "tier"). Populated from the DB on the
# API-key path and from the (future) "plan" claim on the JWT path.
tier: str | None = field(default=None)
# The JWT "plan" claim: a hint for the fail mode only, never authority.
# None on the API-key path; the resolver reads the owner's plan by id.
plan_claim: str | None = field(default=None)


async def get_current_user(
Expand Down Expand Up @@ -189,7 +188,6 @@ async def get_current_user(
# The owning UserDoc is already fetched above for
# email_verified — no extra DB hit to carry the email.
email=user.email.lower() if user and user.email else None,
tier=user.plan.value if user and user.plan else None,
)

# ── JWT path ──────────────────────────────────────────────────────────────
Expand Down Expand Up @@ -246,9 +244,7 @@ async def get_current_user(
email_verified=email_verified,
amr=amr,
email=email,
# Not issued yet — the paid-plans launch adds the claim; TIER
# flag rollouts become a pure data change at that point.
tier=claims.get("plan"),
plan_claim=claims.get("plan"),
scopes=scopes,
app_id=claims.get("app_id"),
)
Expand Down
38 changes: 38 additions & 0 deletions dependencies/entitlements.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
"""
``Entitled``: the resolved entitlements of the request's principal.

Runs the resolver once per request (cache hit in the common case) and hands
the map to routes and services. Nothing on the request path reads
``subscriptions`` or overrides directly.
"""

from __future__ import annotations

from typing import Annotated

from fastapi import Depends, Request

from dependencies.auth import CurrentUser, get_current_user
from services.entitlements import EntitlementService, Resolved


def get_entitlement_service(request: Request) -> EntitlementService:
return request.app.state.entitlement_service


async def get_entitlements(
request: Request,
user: CurrentUser | None = Depends(get_current_user),
service: EntitlementService = Depends(get_entitlement_service),
) -> Resolved:
resolved = await service.resolve_for(
user.user_id if user else None,
plan_hint=user.plan_claim if user else None,
)
if user is not None and not resolved.degraded:
request.state.entitlements_version = resolved.version
return resolved


Entitled = Annotated[Resolved, Depends(get_entitlements)]
EntitlementSvc = Annotated[EntitlementService, Depends(get_entitlement_service)]
69 changes: 68 additions & 1 deletion dependencies/wiring.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from fastapi import FastAPI

from config import AppSettings
from infrastructure.cache.entitlement_cache import EntitlementCache
from infrastructure.cache.feature_flag_cache import FeatureFlagCache
from infrastructure.cache.meta_fetch_cache import MetaFetchCache
from infrastructure.cache.onboarding_cache import OnboardingCache
Expand All @@ -37,6 +38,10 @@
from repositories.blocked_url_repository import BlockedUrlRepository
from repositories.click_repository import ClickRepository
from repositories.custom_domain_repository import CustomDomainRepository
from repositories.entitlement_event_repository import EntitlementEventRepository
from repositories.entitlement_override_repository import (
EntitlementOverrideRepository,
)
from repositories.feature_flag_repository import FeatureFlagRepository
from repositories.feed_domain_repository import FeedDomainRepository
from repositories.legacy.emoji_url_repository import EmojiUrlRepository
Expand All @@ -47,6 +52,7 @@
ReportSubmissionRepository,
)
from repositories.scheduled_task_repository import ScheduledTaskRepository
from repositories.subscription_repository import SubscriptionRepository
from repositories.tag_repository import TagRepository
from repositories.token_repository import TokenRepository
from repositories.url_repository import UrlRepository
Expand Down Expand Up @@ -79,6 +85,7 @@
from services.custom_domain_service import CustomDomainService
from services.domain_intel_service import DomainIntelService
from services.edge_cache.og_writethrough import OgEdgeWritethrough
from services.entitlements import EntitlementService
from services.events.sinks import (
InlineDomainEventSink,
NullDomainEventSink,
Expand All @@ -87,6 +94,7 @@
from services.export.formatters import default_formatters
from services.export.service import ExportService
from services.feature_flag_service import FeatureFlagService
from services.features.catalog import validate_override
from services.meta_tags.sinks import NullMetaImageSink, RedisStreamMetaImageSink
from services.mock_dcv_backend import MockDcvBackend
from services.oauth_service import OAuthService
Expand Down Expand Up @@ -258,6 +266,49 @@ def build_posthog_eraser(settings: AppSettings, http_client) -> PostHogEraser:
)


def build_entitlement_store(
db, redis_client
) -> tuple[
EntitlementCache,
EntitlementEventRepository,
SubscriptionRepository,
EntitlementOverrideRepository,
]:
"""The three entitlement repositories sharing one cache, so every write
to subscriptions or overrides invalidates the same ``ent:{id}`` key."""
cache = EntitlementCache(redis_client)
events = EntitlementEventRepository(db["entitlement_events"])
subscriptions = SubscriptionRepository(db["subscriptions"], events, cache)
overrides = EntitlementOverrideRepository(
db["entitlement_overrides"], events, cache, check=validate_override
)
return cache, events, subscriptions, overrides


def build_entitlement_service(
db, settings: AppSettings, redis_client, *, store=None
) -> EntitlementService:
cache, events, subscriptions, overrides = store or build_entitlement_store(
db, redis_client
)
return EntitlementService(
subscriptions,
overrides,
events,
cache,
selfhost=settings.billing.selfhost,
usage={
"custom_domains_max": CustomDomainRepository(
db["custom_domains"]
).count_by_owner,
"webhook_endpoints_max": WebhookEndpointRepository(
db["webhook-endpoints"]
).count_by_user,
"api_keys_max": ApiKeyRepository(db["api-keys"]).count_by_user,
},
)


def build_account_erasure_service(
db,
settings: AppSettings,
Expand Down Expand Up @@ -355,6 +406,9 @@ def build_account_erasure_service(
redis_client=redis_client,
url_service=url_service,
)
_, ent_events, subscription_repo, override_repo = build_entitlement_store(
db, redis_client
)

return AccountErasureService(
user_repo=user_repo,
Expand All @@ -372,6 +426,9 @@ def build_account_erasure_service(
report_repo=ReportRepository(db["reports"]),
report_submission_repo=ReportSubmissionRepository(db["report_submissions"]),
feature_flag_repo=FeatureFlagRepository(db["feature_flags"]),
subscription_repo=subscription_repo,
override_repo=override_repo,
entitlement_event_repo=ent_events,
r2_storage=r2_storage,
posthog=build_posthog_eraser(settings, http_client),
mailer=build_erasure_mailer(settings, http_client),
Expand Down Expand Up @@ -403,6 +460,11 @@ def wire_services(app: FastAPI, settings: AppSettings, redis_client) -> None:
blocked_url_repo = BlockedUrlRepository(db["blocked-urls"])
app_grant_repo = AppGrantRepository(db["app-grants"])
feature_flag_repo = FeatureFlagRepository(db["feature_flags"])
ent_store = build_entitlement_store(db, redis_client)
_, ent_events, subscription_repo, override_repo = ent_store
app.state.entitlement_service = build_entitlement_service(
db, settings, redis_client, store=ent_store
)

# ── Infrastructure ───────────────────────────────────────────────────
url_cache = UrlCache(redis_client, ttl_seconds=settings.redis.redis_ttl_seconds)
Expand Down Expand Up @@ -773,7 +835,9 @@ def wire_services(app: FastAPI, settings: AppSettings, redis_client) -> None:
max_active_keys=settings.max_active_api_keys,
)
app.state.page_layout_service = PageLayoutService(page_layout_repo)
token_factory = TokenFactory(settings.jwt)
token_factory = TokenFactory(
settings.jwt, plan_of=app.state.entitlement_service.plan_hint_for
)
otp_service = OtpService(token_repo)

app.state.user_repo = user_repo
Expand Down Expand Up @@ -1012,6 +1076,9 @@ def wire_services(app: FastAPI, settings: AppSettings, redis_client) -> None:
report_repo=report_repo,
report_submission_repo=report_submission_repo,
feature_flag_repo=feature_flag_repo,
subscription_repo=subscription_repo,
override_repo=override_repo,
entitlement_event_repo=ent_events,
r2_storage=r2_storage,
posthog=build_posthog_eraser(settings, http_client),
mailer=erasure_mailer,
Expand Down
10 changes: 10 additions & 0 deletions errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,10 @@
"details": NotRequired[Any],
"message": NotRequired[str],
"hint": NotRequired[str],
"feature": NotRequired[str],
"limit": NotRequired[str],
"max": NotRequired[int],
"current": NotRequired[int],
},
)

Expand Down Expand Up @@ -121,6 +125,12 @@ class ConflictError(AppError):
error_code = "conflict"


class InvalidTransitionError(ConflictError):
"""A subscription event arrived in a status it cannot legally change."""

error_code = "invalid_transition"


class BlockedUrlError(AppError):
status_code = 451
error_code = "blocked"
Expand Down
Loading