Stock-Tool/api/main.py

149 lines
5.3 KiB
Python

import os
import hmac
import hashlib
import time
import psycopg2
import psycopg2.extras
from fastapi import FastAPI, HTTPException, Depends, Header
from fastapi.responses import FileResponse
from pydantic import BaseModel
from jose import jwt, JWTError
from datetime import datetime, timedelta, timezone
from typing import Optional
from dotenv import load_dotenv
load_dotenv("/srv/api/.env")
app = FastAPI()
DB_HOST = "127.0.0.1"
DB_NAME = "verimund"
DB_USER = "verimund_user"
DB_PASS = os.getenv("DB_PASS", "Shlevison2k17")
JWT_SECRET = os.getenv("JWT_SECRET", "")
HMAC_SECRET = os.getenv("HMAC_SECRET", "")
JWT_ALGORITHM = "HS256"
JWT_EXPIRE_HOURS = 24
def get_db():
conn = psycopg2.connect(host=DB_HOST, dbname=DB_NAME, user=DB_USER, password=DB_PASS)
try:
yield conn
finally:
conn.close()
def verify_hmac(license_key: str, timestamp: str, signature: str) -> bool:
expected = hmac.new(
HMAC_SECRET.encode(),
f"{license_key}{timestamp}".encode(),
hashlib.sha256
).hexdigest()
return hmac.compare_digest(expected, signature)
def verify_jwt(authorization: Optional[str] = Header(None)):
if not authorization or not authorization.startswith("Bearer "):
raise HTTPException(status_code=401, detail="Missing token")
token = authorization.split(" ")[1]
try:
payload = jwt.decode(token, JWT_SECRET, algorithms=[JWT_ALGORITHM])
return payload
except JWTError:
raise HTTPException(status_code=401, detail="Invalid token")
class AuthRequest(BaseModel):
license_key: str
hwid: str
timestamp: str
signature: str
@app.post("/auth")
def authenticate(req: AuthRequest, db=Depends(get_db)):
try:
req_time = int(req.timestamp)
if abs(time.time() - req_time) > 30:
raise HTTPException(status_code=401, detail="Request expired")
except ValueError:
raise HTTPException(status_code=400, detail="Invalid timestamp")
if not verify_hmac(req.license_key, req.timestamp, req.signature):
raise HTTPException(status_code=401, detail="Invalid signature")
cur = db.cursor(cursor_factory=psycopg2.extras.RealDictCursor)
cur.execute("SELECT * FROM licenses WHERE license_key = %s", (req.license_key,))
license = cur.fetchone()
if not license:
raise HTTPException(status_code=401, detail="Invalid license")
if not license["active"]:
raise HTTPException(status_code=401, detail="License inactive")
if license["expiry_date"] and license["expiry_date"] < datetime.now(timezone.utc):
raise HTTPException(status_code=401, detail="License expired")
if license["hwid"] and license["hwid"] != req.hwid:
raise HTTPException(status_code=401, detail="HWID mismatch")
if not license["hwid"]:
cur.execute("UPDATE licenses SET hwid = %s WHERE license_key = %s", (req.hwid, req.license_key))
db.commit()
payload = {
"license_key": req.license_key,
"hwid": req.hwid,
"channel": license["channel"],
"exp": datetime.now(timezone.utc) + timedelta(hours=JWT_EXPIRE_HOURS)
}
token = jwt.encode(payload, JWT_SECRET, algorithm=JWT_ALGORITHM)
return {"token": token, "channel": license["channel"]}
@app.get("/cache/analyst")
def get_analyst_cache(payload=Depends(verify_jwt), db=Depends(get_db)):
cur = db.cursor(cursor_factory=psycopg2.extras.RealDictCursor)
table = "beta_analyst_cache" if payload.get("channel") == "beta" else "analyst_cache"
cur.execute(f"SELECT * FROM {table}")
return {"data": cur.fetchall()}
@app.get("/cache/fundamentals")
def get_fundamentals_cache(payload=Depends(verify_jwt), db=Depends(get_db)):
cur = db.cursor(cursor_factory=psycopg2.extras.RealDictCursor)
table = "beta_fundamentals_cache" if payload.get("channel") == "beta" else "fundamentals_cache"
cur.execute(f"SELECT * FROM {table}")
return {"data": cur.fetchall()}
@app.get("/cache/sector")
def get_sector_stats(payload=Depends(verify_jwt), db=Depends(get_db)):
cur = db.cursor(cursor_factory=psycopg2.extras.RealDictCursor)
cur.execute("SELECT * FROM sector_stats ORDER BY sector")
return {"data": cur.fetchall()}
@app.get("/update")
def check_update(payload=Depends(verify_jwt), db=Depends(get_db)):
cur = db.cursor(cursor_factory=psycopg2.extras.RealDictCursor)
channel = payload.get("channel", "stable")
cur.execute("SELECT * FROM app_versions WHERE channel = %s ORDER BY created_at DESC LIMIT 1", (channel,))
version = cur.fetchone()
if not version:
raise HTTPException(status_code=404, detail="No version found")
return {"version": version["version"], "channel": channel, "release_notes": version.get("release_notes", "")}
@app.post("/download")
def download_exe(payload=Depends(verify_jwt), db=Depends(get_db)):
cur = db.cursor(cursor_factory=psycopg2.extras.RealDictCursor)
channel = payload.get("channel", "stable")
cur.execute("SELECT * FROM app_versions WHERE channel = %s ORDER BY created_at DESC LIMIT 1", (channel,))
version = cur.fetchone()
if not version:
raise HTTPException(status_code=404, detail="No version found")
path = f"/srv/releases/{channel}/StockScreener-v{version['version']}.exe"
if not os.path.exists(path):
raise HTTPException(status_code=404, detail="File not found")
return FileResponse(path, filename=f"StockScreener-v{version['version']}.exe")