""" Unified AI Proxy — main.py Self-hosted, OpenAI-compatible API gateway + encrypted key manager + dashboard. """ import os import re import json import time import uuid import secrets from contextlib import asynccontextmanager from datetime import datetime from typing import AsyncGenerator import httpx from cryptography.fernet import Fernet from fastapi import FastAPI, Depends, HTTPException, Request from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import HTMLResponse, JSONResponse, StreamingResponse, Response from pydantic import BaseModel from sqlalchemy import ( Column, String, Boolean, Integer, DateTime, Float, Text, create_engine, select, update, delete ) from sqlalchemy.orm import DeclarativeBase, sessionmaker from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker # ── Database & Encryption Storage ───────────────────────────────────────────── if os.path.isdir("/app/data"): ASYNC_DB = "sqlite+aiosqlite:////app/data/db.sqlite" SYNC_DB = "sqlite:////app/data/db.sqlite" elif os.path.isdir("/data"): ASYNC_DB = "sqlite+aiosqlite:////data/db.sqlite" SYNC_DB = "sqlite:////data/db.sqlite" else: ASYNC_DB = "sqlite+aiosqlite:///./db.sqlite" SYNC_DB = "sqlite:///./db.sqlite" class Base(DeclarativeBase): pass class Provider(Base): __tablename__ = "providers" id = Column(Integer, primary_key=True, autoincrement=True) name = Column(String(100), unique=True, nullable=False) base_url = Column(String(500), nullable=False) api_key = Column(String(500), nullable=False) enabled = Column(Boolean, default=True) is_default = Column(Boolean, default=False) created_at = Column(DateTime, default=datetime.utcnow) class RequestLog(Base): __tablename__ = "request_logs" id = Column(Integer, primary_key=True, autoincrement=True) timestamp = Column(DateTime, default=datetime.utcnow) provider_name = Column(String(100)) model = Column(String(200)) status_code = Column(Integer) latency_ms = Column(Float) path = Column(String(200)) streaming = Column(Boolean, default=False) class Settings(Base): __tablename__ = "settings" key = Column(String(100), primary_key=True) value = Column(Text, nullable=False) def init_db(): engine = create_engine(SYNC_DB, connect_args={"check_same_thread": False}) Base.metadata.create_all(engine) Session = sessionmaker(bind=engine) with Session() as s: # Get or generate encryption key enc_row = s.get(Settings, "encryption_key") if not enc_row: enc_key = Fernet.generate_key().decode() s.add(Settings(key="encryption_key", value=enc_key)) else: enc_key = enc_row.value cipher = Fernet(enc_key.encode()) # Master key (check env secret first) env_master = os.environ.get("MASTER_KEY") or os.environ.get("MASTER_API_KEY") or os.environ.get("PROXY_MASTER_KEY") existing = s.get(Settings, "master_key") if env_master: if existing: existing.value = env_master else: s.add(Settings(key="master_key", value=env_master)) print("[startup] Master key loaded securely from Space Secret.") elif not existing: key = f"umk-{secrets.token_urlsafe(32)}" s.add(Settings(key="master_key", value=key)) print(f"\n{'='*60}\n [SECURITY WARNING] Auto-generated MASTER API KEY:\n {key}\n Since this Space is public, set 'MASTER_KEY' in Space Secrets to hide this log!\n{'='*60}\n") else: print(f"[startup] Master key loaded from DB (ends: ...{existing.value[-4:]})") # Auto-seed from env: PROVIDER__URL + PROVIDER__KEY for env_k, env_v in os.environ.items(): m = re.match(r"^PROVIDER_([A-Z0-9_]+)_URL$", env_k) if m: pname = m.group(1).capitalize() key_env = f"PROVIDER_{m.group(1)}_KEY" pkey = os.environ.get(key_env, "") if pkey and not s.query(Provider).filter_by(name=pname).first(): enc_pkey = cipher.encrypt(pkey.encode()).decode() s.add(Provider(name=pname, base_url=env_v, api_key=enc_pkey)) print(f"[startup] Seeded provider: {pname}") s.commit() return cipher def _decrypt(cipher: Fernet, enc: str) -> str: if not enc: return "" try: return cipher.decrypt(enc.encode()).decode() except Exception: return enc # ── Lifespan ────────────────────────────────────────────────────────────────── @asynccontextmanager async def lifespan(app: FastAPI): cipher = init_db() app.state.cipher = cipher engine = create_async_engine(ASYNC_DB) app.state.db = async_sessionmaker(engine, expire_on_commit=False) yield await engine.dispose() app = FastAPI(title="Unified AI Proxy", lifespan=lifespan) app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"]) TIMEOUT = httpx.Timeout(120.0, connect=10.0) # ── Auth & Error Handlers ───────────────────────────────────────────────────── async def verify_key(request: Request): auth = request.headers.get("Authorization", "") api_key_header = request.headers.get("api-key", "") token = "" if auth.startswith("Bearer "): token = auth.removeprefix("Bearer ").strip() elif api_key_header: token = api_key_header.strip() if not token: raise HTTPException(401, "Missing API key in Authorization (Bearer) or api-key header.") async with request.app.state.db() as s: row = await s.get(Settings, "master_key") if not row or not secrets.compare_digest(row.value, token): raise HTTPException(401, "Invalid master API key.") @app.exception_handler(HTTPException) async def openai_exception_handler(request: Request, exc: HTTPException): if request.url.path.startswith("/v1/"): return JSONResponse( status_code=exc.status_code, content={ "error": { "message": str(exc.detail), "type": "invalid_request_error" if exc.status_code < 500 else "api_error", "param": None, "code": exc.status_code } }, headers={"x-request-id": f"req_{uuid.uuid4().hex}"} ) return JSONResponse(status_code=exc.status_code, content={"detail": exc.detail}) @app.get("/health") async def health(request: Request): async with request.app.state.db() as s: result = await s.execute(select(Provider).where(Provider.enabled == True)) n = len(result.scalars().all()) return {"status": "ok", "providers": n} # ── Dashboard ───────────────────────────────────────────────────────────────── DASHBOARD_HTML = r""" Unified AI Proxy
● connecting…
📖 How to use the proxy

Route to a specific provider using providerName/modelName:

curl https://YOUR-SPACE.hf.space/v1/chat/completions \
  -H "Authorization: Bearer YOUR_MASTER_KEY" \
  -H "Content-Type: application/json" \
  -d '{"model":"groq/llama3-8b-8192","messages":[{"role":"user","content":"Hello!"}]}'
Add Provider
NameBase URLKeyDefaultEnabledActions
Loading…

Request Logs

Auto-refreshes every 5s · last 200
TimeProviderModelPathStatusLatencyStream
No requests yet.
Security & Encryption
🔒
Provider API Keys Encrypted at Rest
All backend provider keys are stored encrypted in SQLite (AES-128 via Fernet).
Checking Master Key Security…
Master API Key

Authenticates all proxy requests and dashboard access.

loading…

⚠ Regenerating invalidates the current key immediately.

Default Provider

Requests without a providerName/ prefix route here. Set via the Default toggle in Providers.

loading…
Base Endpoint
""" @app.get("/", response_class=HTMLResponse) async def dashboard(): return DASHBOARD_HTML # ── Provider CRUD ───────────────────────────────────────────────────────────── class ProviderCreate(BaseModel): name: str; base_url: str; api_key: str class ProviderPatch(BaseModel): base_url: str | None = None api_key: str | None = None enabled: bool | None = None is_default: bool | None = None @app.get("/api/providers", dependencies=[Depends(verify_key)]) async def get_providers(request: Request): cipher = request.app.state.cipher async with request.app.state.db() as s: res = await s.execute(select(Provider).order_by(Provider.id)) ps = res.scalars().all() out = [] for p in ps: dec = _decrypt(cipher, p.api_key) out.append({ "id": p.id, "name": p.name, "base_url": p.base_url, "key_preview": f"...{dec[-4:]}" if len(dec)>=4 else "****", "enabled": p.enabled, "is_default": p.is_default }) return out @app.post("/api/providers", dependencies=[Depends(verify_key)]) async def create_provider(body: ProviderCreate, request: Request): cipher = request.app.state.cipher enc_key = cipher.encrypt(body.api_key.encode()).decode() async with request.app.state.db() as s: ex = await s.execute(select(Provider).where(Provider.name == body.name)) if ex.scalar(): raise HTTPException(400, f"Provider '{body.name}' already exists") p = Provider(name=body.name, base_url=body.base_url, api_key=enc_key) s.add(p); await s.commit(); await s.refresh(p) return {"id": p.id, "name": p.name} @app.patch("/api/providers/{pid}", dependencies=[Depends(verify_key)]) async def patch_provider(pid: int, body: ProviderPatch, request: Request): cipher = request.app.state.cipher async with request.app.state.db() as s: p = await s.get(Provider, pid) if not p: raise HTTPException(404, "Not found") if body.base_url is not None: p.base_url = body.base_url if body.api_key is not None: p.api_key = cipher.encrypt(body.api_key.encode()).decode() if body.enabled is not None: p.enabled = body.enabled if body.is_default is not None: if body.is_default: await s.execute(update(Provider).where(Provider.id != pid).values(is_default=False)) p.is_default = body.is_default await s.commit() return {"ok": True} @app.delete("/api/providers/{pid}", dependencies=[Depends(verify_key)]) async def delete_provider(pid: int, request: Request): async with request.app.state.db() as s: p = await s.get(Provider, pid) if not p: raise HTTPException(404, "Not found") await s.delete(p); await s.commit() return {"ok": True} # ── Logs ────────────────────────────────────────────────────────────────────── @app.get("/api/logs", dependencies=[Depends(verify_key)]) async def get_logs(request: Request): async with request.app.state.db() as s: res = await s.execute(select(RequestLog).order_by(RequestLog.id.desc()).limit(200)) logs = res.scalars().all() return [{"id":l.id,"timestamp":l.timestamp.isoformat() if l.timestamp else None, "provider_name":l.provider_name,"model":l.model,"status_code":l.status_code, "latency_ms":round(l.latency_ms,1) if l.latency_ms else None, "path":l.path,"streaming":l.streaming} for l in logs] # ── Settings ────────────────────────────────────────────────────────────────── @app.get("/api/settings/key", dependencies=[Depends(verify_key)]) async def get_key(request: Request): async with request.app.state.db() as s: row = await s.get(Settings, "master_key") k = row.value if row else "" is_env = bool(os.environ.get("MASTER_KEY") or os.environ.get("MASTER_API_KEY") or os.environ.get("PROXY_MASTER_KEY")) return { "key_preview": f"...{k[-8:]}" if len(k)>=8 else "****", "full_key": k, "is_env": is_env } @app.post("/api/settings/key/regenerate", dependencies=[Depends(verify_key)]) async def regen_key(request: Request): new_key = f"umk-{secrets.token_urlsafe(32)}" async with request.app.state.db() as s: row = await s.get(Settings, "master_key") if row: row.value = new_key else: s.add(Settings(key="master_key", value=new_key)) await s.commit() return {"key": new_key} # ── Proxy helpers ───────────────────────────────────────────────────────────── async def _log(session, provider_name, model, status_code, latency_ms, path, streaming): session.add(RequestLog(timestamp=datetime.utcnow(), provider_name=provider_name, model=model, status_code=status_code, latency_ms=latency_ms, path=path, streaming=streaming)) await session.commit() # Prune to 500 rows res = await session.execute(select(RequestLog.id).order_by(RequestLog.id.desc()).offset(500)) old = res.scalars().all() if old: await session.execute(delete(RequestLog).where(RequestLog.id.in_(old))) await session.commit() def _parse_provider(model: str, providers): if "/" in model: prefix, real = model.split("/", 1) for p in providers: if p.name.lower() == prefix.lower(): return p, real return None, model # ── Models aggregation & retrieval (OpenAI Spec) ────────────────────────────── @app.get("/v1/models/{model_id:path}", dependencies=[Depends(verify_key)]) async def retrieve_model(request: Request, model_id: str): async with request.app.state.db() as s: res = await s.execute(select(Provider).where(Provider.enabled == True)) providers = res.scalars().all() provider, _ = _parse_provider(model_id, providers) pname = provider.name if provider else "proxy" return JSONResponse( { "id": model_id, "object": "model", "created": int(time.time()), "owned_by": pname }, headers={"x-request-id": f"req_{uuid.uuid4().hex}"} ) @app.get("/v1/models", dependencies=[Depends(verify_key)]) async def list_models(request: Request): async with request.app.state.db() as s: res = await s.execute(select(Provider).where(Provider.enabled == True)) providers = res.scalars().all() all_models = [] async with httpx.AsyncClient(timeout=httpx.Timeout(10.0)) as client: for p in providers: dec_key = _decrypt(request.app.state.cipher, p.api_key) try: r = await client.get(f"{p.base_url.rstrip('/')}/models", headers={"Authorization": f"Bearer {dec_key}"}) if r.status_code == 200: for m in r.json().get("data", []): m["id"] = f"{p.name.lower()}/{m['id']}" all_models.append(m) except Exception: pass return JSONResponse( {"object": "list", "data": all_models}, headers={"x-request-id": f"req_{uuid.uuid4().hex}"} ) # ── Generic proxy ───────────────────────────────────────────────────────────── @app.api_route("/v1/{path:path}", methods=["GET","POST","PUT","DELETE","PATCH"], dependencies=[Depends(verify_key)]) async def proxy(request: Request, path: str): body_bytes = await request.body() body_json = None model_raw = "" if body_bytes: try: body_json = json.loads(body_bytes) model_raw = body_json.get("model", "") except Exception: pass async with request.app.state.db() as s: res = await s.execute(select(Provider).where(Provider.enabled == True)) providers = res.scalars().all() if not providers: raise HTTPException(503, "No enabled AI providers configured.") provider, real_model = _parse_provider(model_raw, providers) if provider is None: for p in providers: if p.is_default: provider = p; break if provider is None: provider = providers[0] if body_json is not None and model_raw: body_json["model"] = real_model body_bytes = json.dumps(body_json).encode() dec_api_key = _decrypt(request.app.state.cipher, provider.api_key) is_streaming = bool(body_json and body_json.get("stream")) target_url = f"{provider.base_url.rstrip('/')}/{path.lstrip('/')}" req_id = f"req_{uuid.uuid4().hex}" skip_headers = {"host", "authorization", "api-key", "content-length"} forward_headers = {k: v for k, v in request.headers.items() if k.lower() not in skip_headers} forward_headers["Authorization"] = f"Bearer {dec_api_key}" forward_headers["Content-Type"] = request.headers.get("Content-Type", "application/json") start = time.time() if is_streaming: logged = False async def generate() -> AsyncGenerator[bytes, None]: nonlocal logged try: async with httpx.AsyncClient(timeout=TIMEOUT) as client: async with client.stream(request.method, target_url, headers=forward_headers, content=body_bytes) as resp: latency = (time.time() - start) * 1000 if not logged: logged = True async with request.app.state.db() as ls: await _log(ls, provider.name, real_model, resp.status_code, latency, path, True) async for chunk in resp.aiter_bytes(): yield chunk except Exception as e: err_chunk = json.dumps({"error": {"message": str(e), "type": "api_error", "code": 502}}).encode() yield b"data: " + err_chunk + b"\n\n" yield b"data: [DONE]\n\n" return StreamingResponse( generate(), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no", "x-request-id": req_id} ) try: async with httpx.AsyncClient(timeout=TIMEOUT) as client: resp = await client.request(request.method, target_url, headers=forward_headers, content=body_bytes) latency = (time.time() - start) * 1000 async with request.app.state.db() as s: await _log(s, provider.name, real_model, resp.status_code, latency, path, False) resp_headers = {"x-request-id": req_id} content_type = resp.headers.get("content-type", "application/json") if "json" in content_type: return JSONResponse(content=resp.json(), status_code=resp.status_code, headers=resp_headers) return Response(content=resp.content, status_code=resp.status_code, media_type=content_type, headers=resp_headers) except httpx.TimeoutException: raise HTTPException(504, "Upstream AI provider connection timed out.") except Exception as e: raise HTTPException(502, f"Proxy error communicating with upstream provider: {e}")