diff --git a/.gitignore b/.gitignore index 23c3d2b2..4a594399 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,3 @@ __pycache__ .env -active_keys.json \ No newline at end of file +keys.db diff --git a/active_keys.json b/active_keys.json new file mode 100644 index 00000000..2addab23 --- /dev/null +++ b/active_keys.json @@ -0,0 +1,3 @@ +{ + "43492270c8add0bd0c9914ed2811aced168ac45bc3de00d51013e8b9e6b72825": 952 +} \ No newline at end of file diff --git a/proxy/auth.py b/proxy/auth.py index 2e4ca7dc..c4610d8d 100644 --- a/proxy/auth.py +++ b/proxy/auth.py @@ -1,13 +1,9 @@ -import json -import os import hashlib -from fastapi import HTTPException -from .redeem import redeem -ACTIVE_KEYS_FILE = "active_keys.json" -# These are needed here because validate_api_key and pay_for_request use them -RECIEIVE_LN_ADDRESS = os.environ["RECIEIVE_LN_ADDRESS"] -COST_PER_REQUEST = int(os.environ["COST_PER_REQUEST"]) +from fastapi import HTTPException + +from .redeem import redeem +from .db import ApiKey, create_session, RECIEIVE_LN_ADDRESS, COST_PER_REQUEST def _hash_api_key(api_key: str) -> str: @@ -15,79 +11,72 @@ def _hash_api_key(api_key: str) -> str: return hashlib.sha256(api_key.encode()).hexdigest() -def _load_active_keys() -> dict[str, int]: - """Loads the active keys (hashed) and their balances from the JSON file.""" - try: - with open(ACTIVE_KEYS_FILE, "r") as f: - return json.load(f) - except (FileNotFoundError, json.JSONDecodeError): - # Return empty dict if file not found or empty/invalid JSON - return {} - - -def _save_active_keys(active_keys: dict[str, int]) -> None: - """Saves the active keys (hashed) and their balances to the JSON file.""" - # Create directory if it doesn't exist - os.makedirs(os.path.dirname(ACTIVE_KEYS_FILE) or ".", exist_ok=True) - with open(ACTIVE_KEYS_FILE, "w") as f: - json.dump(active_keys, f, indent=2) # Add indent for readability - - async def validate_api_key(api_key: str) -> None: """ - Validates the provided API key. + Validates the provided API key using SQLModel. If it's a cashu key, it redeems it and stores its hash and balance. Otherwise checks if the hash of the key exists. """ if not api_key: - raise HTTPException(status_code=401, detail="API key required") + raise HTTPException(status_code=401, detail="api-key or cashu-token required") hashed_key = _hash_api_key(api_key) - active_keys = _load_active_keys() - if hashed_key in active_keys: - # Key hash already exists, it's valid (might be cashu or other type) - return - - # If hash not found, check if it's a potentially new cashu key - if api_key.startswith("cashu"): - try: - print(f"Attempting to redeem cashu key: {api_key[:15]}...{api_key[-15:]}") - # Redeem the original cashu key - amount = await redeem(api_key, RECIEIVE_LN_ADDRESS) - print(f"Redeemed successfully. Amount: {amount}") - # Store the hash and the redeemed amount - active_keys[hashed_key] = amount - _save_active_keys(active_keys) + async with create_session() as session: + # check if key exists + if await session.get(ApiKey, hashed_key): return - except Exception as e: - print(f"Redemption failed: {e}") - # Include the redemption error message for better debugging - raise HTTPException( - status_code=401, detail=f"Invalid or expired cashu key: {e}" - ) - # If it's not a known hash and not a valid new cashu key - raise HTTPException(status_code=401, detail="Invalid API key") + # If hash not found, check if it's a potentially new cashu key + if api_key.startswith("cashu"): + try: + print( + f"Attempting to redeem cashu key: {api_key[:15]}...{api_key[-15:]}" + ) + # Redeem the original cashu key + amount = await redeem(api_key, RECIEIVE_LN_ADDRESS) + print(f"Redeemed successfully. Amount: {amount}") + # Store the hash and the redeemed amount using SQLModel + new_key = ApiKey(hashed_key=hashed_key, balance=amount) + session.add(new_key) + await session.commit() + await session.refresh(new_key) + return + except Exception as e: + print(f"Redemption failed: {e}") + # Include the redemption error message for better debugging + raise HTTPException( + status_code=401, detail=f"Invalid or expired cashu key: {e}" + ) + + # If it's not a known hash and not a valid new cashu key + raise HTTPException(status_code=401, detail="Invalid API key") async def pay_for_request(api_key: str) -> None: - """Deducts the cost of a request from the balance associated with the API key hash.""" + """Deducts the cost of a request from the balance associated with the API key hash using SQLModel.""" hashed_key = _hash_api_key(api_key) - active_keys = _load_active_keys() - if hashed_key not in active_keys: - # This should theoretically not happen if validate_api_key was called first - raise HTTPException(status_code=401, detail="API key not validated") + # Get the key record using SQLModel + async with create_session() as session: + key_record = await session.get(ApiKey, hashed_key) - if active_keys[hashed_key] < COST_PER_REQUEST: - raise HTTPException( - status_code=402, detail="Insufficient balance" - ) # 402 Payment Required + if key_record is None: + # This should theoretically not happen if validate_api_key was called first + # Consider adding a check or relying on validate_api_key structure + raise HTTPException(status_code=401, detail="API key not validated") - # todo: COST_PER_INPUT_TOKENS + COST_PER_OUTPUT_TOKENS (like openai) - active_keys[hashed_key] -= COST_PER_REQUEST - _save_active_keys(active_keys) - print( - f"Charged {COST_PER_REQUEST}. New balance for key hash {hashed_key[:10]}...: {active_keys[hashed_key]}" - ) + if key_record.balance < COST_PER_REQUEST: + raise HTTPException( + status_code=402, detail="Insufficient balance" + ) # 402 Payment Required + + # todo: COST_PER_INPUT_TOKENS + COST_PER_OUTPUT_TOKENS (like openai) + key_record.balance -= COST_PER_REQUEST + session.add(key_record) # Mark the object as changed + await session.commit() + await session.refresh(key_record) + + print( + f"Charged {COST_PER_REQUEST}. New balance for key hash {hashed_key[:10]}...: {key_record.balance}" + ) diff --git a/proxy/db.py b/proxy/db.py new file mode 100644 index 00000000..3f831b8e --- /dev/null +++ b/proxy/db.py @@ -0,0 +1,42 @@ +from contextlib import asynccontextmanager +import os +from typing import AsyncGenerator +from sqlmodel import Field, SQLModel +from sqlalchemy.ext.asyncio.engine import create_async_engine +from sqlmodel.ext.asyncio.session import AsyncSession + + +DATABASE_URL = os.environ.get("DATABASE_URL", "sqlite+aiosqlite:///keys.db") + +# These are needed here because validate_api_key and pay_for_request use them in auth.py +# Consider passing these as arguments or managing config differently if needed +RECIEIVE_LN_ADDRESS = os.environ["RECIEIVE_LN_ADDRESS"] +COST_PER_REQUEST = int(os.environ["COST_PER_REQUEST"]) + +engine = create_async_engine(DATABASE_URL, echo=False) # echo=True for debugging SQL + + +class ApiKey(SQLModel, table=True): # type: ignore + __tablename__ = "api_keys" + + hashed_key: str = Field(primary_key=True) + balance: int = Field(default=0) + total_spent: int = Field(default=0) + total_requests: int = Field(default=0) + + +async def init_db() -> None: + """Initializes the database and creates tables if they don't exist.""" + async with engine.begin() as conn: + await conn.run_sync(SQLModel.metadata.create_all) + + +async def get_session() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session + + +@asynccontextmanager +async def create_session() -> AsyncGenerator[AsyncSession, None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + yield session