from flask import Flask, request, jsonify, send_from_directory import requests import os app = Flask(__name__, static_folder="static") # Put your *.gradio.live or *.hf.space URL (no trailing slash) GRADIO_SERVER_URL = os.getenv("GRADIO_SERVER_URL", "https://c9ceb67adbae7238f6.gradio.live").rstrip("/") # Default endpoint — some Spaces expose /predict by default GRADIO_PREDICT_URL = f"{GRADIO_SERVER_URL}/predict" @app.route("/") def index(): return send_from_directory(".", "index.html") @app.route("/static/") def static_files(filename): return send_from_directory(app.static_folder, filename) @app.route("/generate", methods=["POST"]) def generate(): body = request.get_json(silent=True) or {} prompt = str(body.get("prompt", "")).strip() if not prompt: return jsonify({"error": "prompt is required"}), 400 # You can optionally override if the Space uses a custom API name. # Example: body["api_name"] = "/generate" or "/my_api" api_name = body.get("api_name", "/predict") if not api_name.startswith("/"): api_name = "/" + api_name api_url = f"{GRADIO_SERVER_URL}{api_name}" # Gradio expects positional arguments if the endpoint is /predict # Usually it's a list matching the input components gradio_payload = { "data": [ prompt, int(body.get("max_new_tokens", 200)), float(body.get("temperature", 0.9)), float(body.get("top_p", 0.95)), float(body.get("repetition_penalty", 1.2)), ] } try: resp = requests.post(api_url, json=gradio_payload, timeout=120) resp.raise_for_status() except requests.exceptions.ConnectionError: return jsonify({"error": "Cannot reach inference server."}), 502 except requests.exceptions.Timeout: return jsonify({"error": "Inference server timed out."}), 504 except requests.exceptions.HTTPError as e: return jsonify({"error": f"Inference server error: {str(e)}", "response": resp.text}), 502 try: result = resp.json() # for many Spaces, the first item in "data" is the output generated_text = result.get("data", [None])[0] if isinstance(result, dict) else None except Exception: return jsonify({"error": "Unexpected inference server response", "raw": resp.text}), 502 return jsonify({"generated_text": generated_text, "prompt": prompt}) if __name__ == "__main__": port = int(os.getenv("PORT", "7860")) app.run(port=port, debug=False)