use sqlite as db

This commit is contained in:
Shroominic
2025-04-21 12:55:07 +08:00
parent 945797d015
commit 3ed834ef1e
4 changed files with 101 additions and 67 deletions
+1 -1
View File
@@ -1,3 +1,3 @@
__pycache__
.env
active_keys.json
keys.db
+3
View File
@@ -0,0 +1,3 @@
{
"43492270c8add0bd0c9914ed2811aced168ac45bc3de00d51013e8b9e6b72825": 952
}
+55 -66
View File
@@ -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}"
)
+42
View File
@@ -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