""" DocuLens — Usage tracking & tier enforcement. Middleware that checks a user's monthly page quota before allowing extraction, and increments the counter after successful processing. """ import os import logging from datetime import date, datetime from typing import Optional logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- # Tier configuration # --------------------------------------------------------------------------- TIER_LIMITS = { "free": 50, "starter": 500, "pro": 2000, "enterprise": 100_000, # effectively unlimited; real cap set per-contract } TIER_FEATURES = { "free": { "pages_per_month": 50, "batch_upload": False, "max_batch_files": 0, "webhooks": False, "priority_models": False, "api_access": False, "support": "community", }, "starter": { "pages_per_month": 500, "batch_upload": True, "max_batch_files": 10, "webhooks": True, "priority_models": False, "api_access": True, "support": "email", }, "pro": { "pages_per_month": 2000, "batch_upload": True, "max_batch_files": 20, "webhooks": True, "priority_models": True, "api_access": True, "support": "priority", }, "enterprise": { "pages_per_month": 100_000, "batch_upload": True, "max_batch_files": 50, "webhooks": True, "priority_models": True, "api_access": True, "support": "dedicated", }, } def _current_period() -> str: """Return the first day of the current month as YYYY-MM-DD.""" today = date.today() return today.replace(day=1).isoformat() # --------------------------------------------------------------------------- # Supabase-backed usage tracking # --------------------------------------------------------------------------- def _get_supabase(): """Return the Supabase client, or None if not configured.""" try: from db.supabase import get_client return get_client() except Exception: return None def get_user_tier(user_id: str) -> str: """Look up the user's tier from user_profiles. Defaults to 'free'.""" sb = _get_supabase() if not sb: return "free" try: resp = sb.table("user_profiles").select("tier").eq("id", user_id).single().execute() if resp.data: return resp.data.get("tier", "free") except Exception as e: logger.warning("Failed to fetch user tier: %s", e) return "free" def get_usage(user_id: str) -> dict: """ Get the user's current month usage and limits. Returns: {tier, pages_used, pages_limit, period_start, remaining} """ tier = get_user_tier(user_id) limit = TIER_LIMITS.get(tier, 50) period = _current_period() sb = _get_supabase() pages_used = 0 if sb: try: resp = ( sb.table("usage_tracking") .select("pages_used") .eq("user_id", user_id) .eq("period_start", period) .single() .execute() ) if resp.data: pages_used = resp.data.get("pages_used", 0) except Exception as e: logger.warning("Failed to fetch usage: %s", e) return { "tier": tier, "pages_used": pages_used, "pages_limit": limit, "period_start": period, "remaining": max(0, limit - pages_used), } def check_quota(user_id: str, pages_requested: int = 1) -> dict: """ Check if the user has enough quota for the requested pages. Returns: {allowed: bool, usage: {...}, message: str} """ usage = get_usage(user_id) allowed = usage["remaining"] >= pages_requested if not allowed: msg = ( f"Monthly quota exceeded. You've used {usage['pages_used']} of " f"{usage['pages_limit']} pages on the {usage['tier']} plan. " f"Upgrade your plan or wait until next month." ) else: msg = "ok" return {"allowed": allowed, "usage": usage, "message": msg} def increment_usage(user_id: str, pages: int = 1) -> bool: """ Increment the user's page count for the current month. Creates the usage row if it doesn't exist (upsert). Returns True on success. """ sb = _get_supabase() if not sb: return False period = _current_period() tier = get_user_tier(user_id) try: # Try to upsert — on conflict (user_id, period_start), increment resp = sb.rpc("increment_usage", { "p_user_id": user_id, "p_period": period, "p_pages": pages, "p_tier": tier, }).execute() return True except Exception: # Fallback: manual upsert if RPC not available try: existing = ( sb.table("usage_tracking") .select("id, pages_used") .eq("user_id", user_id) .eq("period_start", period) .maybe_single() .execute() ) if existing.data: new_count = existing.data["pages_used"] + pages sb.table("usage_tracking").update({ "pages_used": new_count, "updated_at": datetime.utcnow().isoformat(), }).eq("id", existing.data["id"]).execute() else: sb.table("usage_tracking").insert({ "user_id": user_id, "period_start": period, "pages_used": pages, "tier": tier, }).execute() return True except Exception as e: logger.error("Failed to increment usage: %s", e) return False def check_feature(user_id: str, feature: str) -> bool: """Check if a user's tier includes a specific feature.""" tier = get_user_tier(user_id) tier_features = TIER_FEATURES.get(tier, TIER_FEATURES["free"]) return bool(tier_features.get(feature, False)) # --------------------------------------------------------------------------- # SQL function for atomic increment (run in Supabase SQL Editor) # --------------------------------------------------------------------------- USAGE_INCREMENT_SQL = """ -- Atomic usage increment — add to schema_additions.sql and run once CREATE OR REPLACE FUNCTION increment_usage( p_user_id TEXT, p_period DATE, p_pages INTEGER, p_tier TEXT DEFAULT 'free' ) RETURNS void LANGUAGE plpgsql AS $$ BEGIN INSERT INTO usage_tracking (user_id, period_start, pages_used, tier) VALUES (p_user_id, p_period, p_pages, p_tier) ON CONFLICT (user_id, period_start) DO UPDATE SET pages_used = usage_tracking.pages_used + p_pages, updated_at = now(); END; $$; """