mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-06 12:38:23 +00:00
use sqlite as db
This commit is contained in:
+1
-1
@@ -1,3 +1,3 @@
|
||||
__pycache__
|
||||
.env
|
||||
active_keys.json
|
||||
keys.db
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
{
|
||||
"43492270c8add0bd0c9914ed2811aced168ac45bc3de00d51013e8b9e6b72825": 952
|
||||
}
|
||||
+55
-66
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user