File size: 7,007 Bytes
04c4194
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
"""
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;
$$;
"""