#!/usr/bin/env python # -*- coding: utf-8 -*- """ OmniParse AI – B2B SaaS invoice extractor ================================================ * Single‑file Gradio app (ready for Hugging Face Spaces) * SQLite persistence, username/password auth (PBKDF2) * Subscription tier handling (Free / Basic / Pro / Enterprise) * Stripe Checkout + in‑Space webhook (FastAPI) * Mock PDF parser (cpu‑only) – swap with Donut model for GPU later * Premium dark UI built with Tailwind CSS (no external CSS files) * All secrets (Stripe keys) read from environment variables """ import os, json, base64, hashlib, secrets, datetime as dt from pathlib import Path from typing import Dict, Any, Optional import gradio as gr import stripe # Updated import to include FileResponse from fastapi import FastAPI, Request, HTTPException, FileResponse # ---------------------------------------------------------------------- # Gradio hidden backend (no visible UI components) # ---------------------------------------------------------------------- with gr.Blocks() as demo: # Legacy UI block removed pass # Mount custom static UI at root path demo.mount_static("/", path="hf_upload/static") # ---------------------------------------------------------------------- # FastAPI routes for serving the custom UI and API endpoints # ---------------------------------------------------------------------- @fastapi_app.get("/") async def serve_ui(): return FileResponse(path="hf_upload/static/index.html", media_type="text/html") # Placeholder API endpoints – can be expanded to call the existing logic @fastapi_app.post("/api/login") async def api_login(request: Request): data = await request.json() username = data.get("username") password = data.get("password") msg, token, plan, user = login_user(username, password) return JSONResponse({"message": msg, "token": token, "plan": plan, "user": user}) @fastapi_app.post("/api/register") async def api_register(request: Request): data = await request.json() username = data.get("username") email = data.get("email") password = data.get("password") terms = data.get("terms", False) msg, _ = register_user(username, email, password, terms) return JSONResponse({"message": msg}) @fastapi_app.post("/api/upload_invoice") async def api_upload_invoice(request: Request): token = request.headers.get("Authorization", "").replace("Bearer ", "") form = await request.form() file = form.get("file") if not file: return JSONResponse({"error": "No file uploaded"}, status_code=400) # Save temporary file tmp_path = Path("tmp_uploads") / file.filename tmp_path.parent.mkdir(exist_ok=True) with open(tmp_path, "wb") as f: f.write(await file.read()) # Create a simple file-like object for upload_invoice class SimpleFile: def __init__(self, name, path): self.name = name self._path = path def read(self): return open(self._path, "rb").read() simple_file = SimpleFile(name=file.filename, path=tmp_path) status_msg, usage_html, result_html, _ = upload_invoice(simple_file, token, "free") return JSONResponse({"status": status_msg, "usage": usage_html, "result": result_html}) from fastapi.responses import JSONResponse from dotenv import load_dotenv from pdfminer.high_level import extract_text import sqlite3 # ---------------------------------------------------------------------- # Load environment variables (Stripe keys will be added later in the Space) # ---------------------------------------------------------------------- load_dotenv() STRIPE_PUBLISHABLE_KEY = os.getenv("STRIPE_PUBLISHABLE_KEY", "pk_test_placeholder") STRIPE_SECRET_KEY = os.getenv("STRIPE_SECRET_KEY", "sk_test_placeholder") stripe.api_key = STRIPE_SECRET_KEY # ---------------------------------------------------------------------- # Database helpers (SQLite) # ---------------------------------------------------------------------- DB_PATH = Path("omniparse.db") CONN = None def get_conn(): global CONN if CONN is None: CONN = sqlite3.connect(DB_PATH, check_same_thread=False) CONN.row_factory = sqlite3.Row return CONN def init_db(): conn = get_conn() cur = conn.cursor() cur.execute(""" CREATE TABLE IF NOT EXISTS users ( id INTEGER PRIMARY KEY AUTOINCREMENT, username TEXT UNIQUE NOT NULL, email TEXT, password_hash TEXT NOT NULL, salt TEXT NOT NULL, plan TEXT NOT NULL DEFAULT 'free', stripe_customer_id TEXT, accepted_terms INTEGER NOT NULL DEFAULT 0, created_at TEXT NOT NULL ) """) cur.execute(""" CREATE TABLE IF NOT EXISTS invoices ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, filename TEXT NOT NULL, uploaded_at TEXT NOT NULL, status TEXT NOT NULL, parsed_json TEXT, file_hash TEXT, FOREIGN KEY(user_id) REFERENCES users(id) ) """) cur.execute(""" CREATE TABLE IF NOT EXISTS usage ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, month TEXT NOT NULL, count INTEGER NOT NULL, UNIQUE(user_id, month), FOREIGN KEY(user_id) REFERENCES users(id) ) """) conn.commit() # ---------------------------------------------------------------------- # Crypto helpers (PBKDF2 password hashing) # ---------------------------------------------------------------------- def hash_password(password: str, salt: bytes) -> str: dk = hashlib.pbkdf2_hmac('sha256', password.encode('utf-8'), salt, 200_000, dklen=32) return base64.b64encode(dk).decode('utf-8') def verify_password(stored_hash: str, password: str, salt: bytes) -> bool: return stored_hash == hash_password(password, salt) # ---------------------------------------------------------------------- # Subscription limits (per plan) # ---------------------------------------------------------------------- PLAN_LIMITS = {"free": 20, "basic": 200, "pro": 2000, "enterprise": None} PLAN_LABELS = { "free": "Free – 20 invoices / month", "basic": "Basic – 200 invoices / month", "pro": "Pro – 2 000 invoices / month", "enterprise": "Enterprise – Unlimited", } # ---------------------------------------------------------------------- # Stripe price IDs (replace with real IDs in the Space) # ---------------------------------------------------------------------- STRIPE_PRICE_IDS = { "basic": "price_1BasicPlan", "pro": "price_1ProPlan", "enterprise": "price_1EnterprisePlan", } # ---------------------------------------------------------------------- # Mock PDF parser (CPU‑only) # ---------------------------------------------------------------------- def mock_parse_pdf(file_path: Path) -> Dict[str, Any]: text = extract_text(str(file_path)) lines = [ln.strip() for ln in text.splitlines() if ln.strip()] totals = [ln for ln in lines if "TOTAL" in ln.upper()] vendor = lines[0] if lines else "Unknown Vendor" return {"vendor": vendor, "raw_text": text, "totals": totals[:1], "line_items": lines[1:5]} # ---------------------------------------------------------------------- # Helper: current month key (YYYY‑MM) # ---------------------------------------------------------------------- def current_month_key() -> str: return dt.datetime.utcnow().strftime("%Y-%m") # ---------------------------------------------------------------------- # FastAPI app (Stripe webhook) – mounted inside Gradio # ---------------------------------------------------------------------- fastapi_app = FastAPI() @fastapi_app.post("/webhook") async def stripe_webhook(request: Request): payload = await request.body() sig_header = request.headers.get("stripe-signature") endpoint_secret = os.getenv("STRIPE_WEBHOOK_SECRET", "") try: event = stripe.Webhook.construct_event(payload, sig_header, endpoint_secret) except Exception as e: raise HTTPException(status_code=400, detail=str(e)) if event["type"] == "checkout.session.completed": session = event["data"]["object"] customer_id = session.get("customer") price_id = session["display_items"][0]["price"]["id"] plan = next((k for k, v in STRIPE_PRICE_IDS.items() if v == price_id), None) if plan: conn = get_conn() cur = conn.cursor() cur.execute( "UPDATE users SET plan = ?, stripe_customer_id = ? WHERE id = (SELECT user_id FROM usage WHERE month = ? ORDER BY id DESC LIMIT 1)", (plan, customer_id, current_month_key()), ) conn.commit() return JSONResponse({"status": "ok"}) # ---------------------------------------------------------------------- # UI helper functions # ---------------------------------------------------------------------- def render_navbar(username: Optional[str], plan: str) -> str: label = PLAN_LABELS.get(plan, "Free") return f""" """ def render_progress_bar(used: int, limit: Optional[int]) -> str: if limit is None: return "
Unlimited
" percent = int((used / limit) * 100) if limit else 0 return f"""
{used} / {limit} invoices used
""" # ---------------------------------------------------------------------- # Gradio callbacks # ---------------------------------------------------------------------- def register_user(username, email, password, terms): if not terms: return "You must accept the Terms & Conditions.", None conn = get_conn(); cur = conn.cursor() cur.execute("SELECT id FROM users WHERE username = ?", (username,)) if cur.fetchone(): return "Username already taken.", None salt = secrets.token_bytes(16) pwd_hash = hash_password(password, salt) cur.execute( "INSERT INTO users (username, email, password_hash, salt, plan, accepted_terms, created_at) VALUES (?,?,?,?,?,1,?)", (username, email, pwd_hash, base64.b64encode(salt).decode(), "free", dt.datetime.utcnow().isoformat()), ) conn.commit() return "Registration successful! You can now log in.", None def login_user(username, password): conn = get_conn(); cur = conn.cursor() cur.execute("SELECT id, password_hash, salt, plan FROM users WHERE username = ?", (username,)) row = cur.fetchone() if not row or not verify_password(row["password_hash"], password, base64.b64decode(row["salt"])): return "Invalid credentials.", None, None, None token = str(row["id"]) return "Login successful.", token, row["plan"], username def logout_user(): return "", None, "free", None def upload_invoice(file_obj, session_token, plan): if not session_token: return "You must be logged in.", "", "", "" conn = get_conn(); cur = conn.cursor() user_id = int(session_token) month = current_month_key() cur.execute("SELECT count FROM usage WHERE user_id = ? AND month = ?", (user_id, month)) row = cur.fetchone() used = row["count"] if row else 0 limit = PLAN_LIMITS[plan] if limit is not None and used >= limit: return f"Quota exceeded for {plan} plan ({limit} invoices/month).", "", "", "" tmp = Path("tmp_uploads") tmp.mkdir(exist_ok=True) file_path = tmp / file_obj.name with open(file_path, "wb") as f: f.write(file_obj.read()) file_hash = hashlib.sha256(file_path.read_bytes()).hexdigest() cur.execute("SELECT id FROM invoices WHERE user_id = ? AND file_hash = ?", (user_id, file_hash)) if cur.fetchone(): return "Duplicate invoice detected.", "", "", "" parsed = mock_parse_pdf(file_path) parsed_json = json.dumps(parsed, ensure_ascii=False, indent=2) cur.execute( "INSERT INTO invoices (user_id, filename, uploaded_at, status, parsed_json, file_hash) VALUES (?,?,?,?,?,?)", (user_id, file_obj.name, dt.datetime.utcnow().isoformat(), "processed", parsed_json, file_hash), ) if row: cur.execute("UPDATE usage SET count = count + 1 WHERE id = ?", (row["id"],)) else: cur.execute("INSERT INTO usage (user_id, month, count) VALUES (?,?,1)", (user_id, month)) conn.commit() usage_html = render_progress_bar(used + 1, limit) result = f"
{parsed_json}
" return "Invoice processed successfully.", usage_html, result, "" def get_stripe_checkout_url(plan, token): if plan == "free": return "" price_id = STRIPE_PRICE_IDS.get(plan) if not price_id: return "" conn = get_conn(); cur = conn.cursor() cur.execute("SELECT email FROM users WHERE id = ?", (int(token),)) email = cur.fetchone()["email"] if cur.fetchone() else None session = stripe.checkout.Session.create( payment_method_types=["card"], line_items=[{"price": price_id, "quantity": 1}], mode="subscription", success_url="https://huggingface.co/spaces/your-username/OmniParseAI?success=true", cancel_url="https://huggingface.co/spaces/your-username/OmniParseAI?canceled=true", client_reference_id=token, customer_email=email, ) return session.url # ---------------------------------------------------------------------- # Gradio UI definition # ---------------------------------------------------------------------- navbar = gr.HTML(render_navbar(None, "free")) with gr.Column(): with gr.Tab("Login / Register"): with gr.Row(): with gr.Column(): login_u = gr.Textbox(label="Username") login_p = gr.Textbox(label="Password", type="password") login_btn = gr.Button("Log in") login_msg = gr.Markdown() with gr.Column(): reg_u = gr.Textbox(label="New username") reg_e = gr.Textbox(label="Email") reg_p = gr.Textbox(label="Password", type="password") reg_t = gr.Checkbox(label="I accept the Terms & Conditions") reg_btn = gr.Button("Register") reg_msg = gr.Markdown() with gr.Tab("Dashboard"): with gr.Row(visible=False) as dash_row: with gr.Column(): usage_html = gr.HTML() file_upload = gr.File(label="Upload PDF invoice") status_msg = gr.Markdown() result_html = gr.HTML() with gr.Column(): account_md = gr.Markdown() logout_btn = gr.Button("Log out") # Callbacks login_btn.click(fn=login_user, inputs=[login_u, login_p], outputs=[login_msg, session_state, user_plan, username_state]) reg_btn.click(fn=register_user, inputs=[reg_u, reg_e, reg_p, reg_t], outputs=[reg_msg, None]) logout_btn.click(fn=logout_user, inputs=[], outputs=[login_msg, session_state, user_plan, username_state]) file_upload.upload(fn=upload_invoice, inputs=[file_upload, session_state, user_plan], outputs=[status_msg, usage_html, result_html, None]) def refresh_nav(_, __, plan): return render_navbar(username_state.value, plan) login_btn.then(refresh_nav, [], navbar) logout_btn.then(refresh_nav, [], navbar) # ---------------------------------------------------------------------- # Mount FastAPI webhook inside Gradio # ---------------------------------------------------------------------- @fastapi_app.get("/upgrade_url") async def upgrade_url(plan: str, request: Request): token = request.headers.get("Authorization", "").replace("Bearer ", "") if not token: return JSONResponse({"error": "not logged in"}) url = get_stripe_checkout_url(plan, token) return JSONResponse({"url": url}) demo.mount_fastapi_app(fastapi_app, path="/") if __name__ == "__main__": init_db() demo.launch(server_name="0.0.0.0", server_port=7860, share=False)