proxy-space-for / app.py
Bc-AI's picture
Update app.py
b9b4453 verified
Raw History Blame Contribute Delete
10.6 kB
# ─────────────────────────────────────────────────────────────────────────────
# Mira Proxy Space β€’ Docker SDK β€’ No GPU needed
#
# Routes:
# POST /v1/chat/completions β†’ mira-1-large (default)
# POST /v1/{model_key}/chat/completions β†’ any registered model
# GET /v1/models β†’ list available models
# GET /health β†’ proxy + upstream health
# GET / β†’ info page
#
# Env secrets (set in Space settings):
# HF_TOKEN β€” used both to call upstream spaces AND to optionally
# gate inbound requests (if AUTH_REQUIRED=true)
# AUTH_REQUIRED β€” set to "true" to require Bearer token on inbound calls
# ─────────────────────────────────────────────────────────────────────────────
from __future__ import annotations
import os
import time
import logging
from functools import wraps
import requests
from flask import Flask, request, Response, jsonify, abort
from flask_cors import CORS
# ── Logging ───────────────────────────────────────────────────────────────────
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(message)s",
)
log = logging.getLogger(__name__)
# ── App ───────────────────────────────────────────────────────────────────────
app = Flask(__name__)
CORS(app, origins=["*"])
# ── Config ────────────────────────────────────────────────────────────────────
HF_TOKEN = os.environ.get("HF_TOKEN", "")
AUTH_REQUIRED = os.environ.get("AUTH_REQUIRED", "false").lower() == "true"
# ── Endpoint registry ─────────────────────────────────────────────────────────
# Each value is the Space root URL (no trailing path).
# The proxy appends /v1/chat/completions when forwarding.
# Update these slugs to match your actual Space URLs.
ENDPOINTS: dict[str, str] = {
"mira-1-large": "https://ml-intern-explorers-inf-end-1.hf.space",
"mira-1-xl": "https://ml-intern-explorers-inf-end-2.hf.space",
}
DEFAULT_MODEL = "mira-1-large"
# Model metadata for /v1/models response (OpenAI-compatible)
MODEL_META = {
"mira-1-large": {"description": "Mira 1 Large", "context_window": 8192},
"mira-1-xl": {"description": "Mira 1 XL", "context_window": 16384},
}
# ── Auth guard (optional) ─────────────────────────────────────────────────────
def require_auth(f):
@wraps(f)
def decorated(*args, **kwargs):
if not AUTH_REQUIRED:
return f(*args, **kwargs)
auth = request.headers.get("Authorization", "")
if not auth.startswith("Bearer ") or auth.split(" ", 1)[1] != HF_TOKEN:
return jsonify({"error": "Unauthorized"}), 401
return f(*args, **kwargs)
return decorated
# ── Upstream headers ──────────────────────────────────────────────────────────
def _upstream_headers() -> dict:
return {
"Authorization": f"Bearer {HF_TOKEN}",
"Content-Type": "application/json",
"Accept": "application/json",
}
# ── Core proxy logic ──────────────────────────────────────────────────────────
def proxy_request(model_key: str, data: dict) -> Response:
base = ENDPOINTS[model_key].rstrip("/")
target = f"{base}/v1/chat/completions"
# Ensure the model field in the payload matches what the Space expects
data.setdefault("model", model_key)
log.info(f"β†’ {model_key} stream={data.get('stream', False)}")
if data.get("stream", False):
# ── Streaming: forward SSE chunks as they arrive ──────────────────────
upstream = requests.post(
target,
json=data,
headers=_upstream_headers(),
stream=True,
timeout=(10, 300), # (connect, read)
)
upstream.raise_for_status()
def generate():
try:
for chunk in upstream.iter_content(
chunk_size=None, decode_unicode=False
):
if chunk:
yield chunk
except Exception as exc:
log.error(f"Streaming error for {model_key}: {exc}")
yield b"data: [DONE]\n\n"
return Response(
generate(),
status=upstream.status_code,
content_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"X-Accel-Buffering":"no", # disable nginx buffering
"Connection": "keep-alive",
},
)
# ── Non-streaming ─────────────────────────────────────────────────────────
upstream = requests.post(
target,
json=data,
headers=_upstream_headers(),
timeout=(10, 60),
)
upstream.raise_for_status()
return jsonify(upstream.json())
# ── Routes ────────────────────────────────────────────────────────────────────
@app.route("/", methods=["GET"])
def index():
"""Human-readable info page."""
return jsonify({
"service": "Mira Proxy",
"version": "1.0.0",
"models": list(ENDPOINTS.keys()),
"default_model": DEFAULT_MODEL,
"endpoints": {
"chat": "POST /v1/chat/completions",
"models": "GET /v1/models",
"health": "GET /health",
},
})
@app.route("/v1/models", methods=["GET"])
@require_auth
def list_models():
"""OpenAI-compatible model listing."""
now = int(time.time())
return jsonify({
"object": "list",
"data": [
{
"id": key,
"object": "model",
"created": now,
"owned_by": "mira",
**MODEL_META.get(key, {}),
}
for key in ENDPOINTS
],
})
@app.route("/v1/chat/completions", methods=["POST"])
@require_auth
def chat_completions_default():
"""
Default route β€” always routes to mira-1-large.
Existing callers need zero changes.
"""
data = request.json or {}
# If caller already specifies a known model in the payload, honour it.
requested = data.get("model", DEFAULT_MODEL)
model_key = requested if requested in ENDPOINTS else DEFAULT_MODEL
try:
return proxy_request(model_key, data)
except requests.HTTPError as e:
log.error(f"Upstream HTTP error: {e}")
return jsonify({"error": str(e)}), e.response.status_code
except requests.Timeout:
return jsonify({"error": "upstream timeout"}), 504
except Exception as e:
log.error(f"Proxy error: {e}")
return jsonify({"error": str(e)}), 502
@app.route("/v1/<model_key>/chat/completions", methods=["POST"])
@require_auth
def chat_completions_by_model(model_key: str):
"""
Per-model route.
e.g. POST /v1/mira-1-xl/chat/completions
"""
if model_key not in ENDPOINTS:
return jsonify({
"error": f"Unknown model '{model_key}'",
"available_models": list(ENDPOINTS.keys()),
}), 404
data = request.json or {}
try:
return proxy_request(model_key, data)
except requests.HTTPError as e:
log.error(f"Upstream HTTP error ({model_key}): {e}")
return jsonify({"error": str(e)}), e.response.status_code
except requests.Timeout:
return jsonify({"error": "upstream timeout"}), 504
except Exception as e:
log.error(f"Proxy error ({model_key}): {e}")
return jsonify({"error": str(e)}), 502
@app.route("/health", methods=["GET"])
def health():
"""
Proxy liveness + optional upstream health checks.
Pass ?upstream=true to also ping each Space's /health endpoint.
"""
result: dict = {
"status": "ok",
"service": "mira-proxy",
"models": list(ENDPOINTS.keys()),
"upstream": {},
}
if request.args.get("upstream", "false").lower() == "true":
for key, base in ENDPOINTS.items():
url = f"{base.rstrip('/')}/health"
try:
r = requests.get(
url,
headers={"Authorization": f"Bearer {HF_TOKEN}"},
timeout=8,
)
result["upstream"][key] = {
"status": "ok" if r.ok else "error",
"http_status": r.status_code,
"body": r.json() if r.ok else r.text[:200],
}
except requests.Timeout:
result["upstream"][key] = {"status": "timeout"}
except Exception as exc:
result["upstream"][key] = {"status": "error", "detail": str(exc)}
return jsonify(result)
# ── Dev entrypoint ────────────────────────────────────────────────────────────
if __name__ == "__main__":
# gunicorn is used in production (see Dockerfile CMD)
app.run(host="0.0.0.0", port=7860, debug=False)