piston-api / main.py
Aman Githala
fix: check both 'expected' and 'output' field names in test cases
267bea2
Raw History Blame Contribute Delete
8.93 kB
import os
import time
import hmac
from fastapi import FastAPI, HTTPException, Depends, Header, Request
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, field_validator
from redis import Redis
import httpx
# --- Configuration ---
EXECUTION_API_KEY = os.getenv("EXECUTION_API_KEY", "")
COMPILER_API_KEY = os.getenv("COMPILER_API_KEY", "")
COMPILER_API_URL = "https://api.onlinecompiler.io/api/run-code-sync/"
SUPABASE_URL = os.getenv("SUPABASE_URL", "")
SUPABASE_SERVICE_ROLE_KEY = os.getenv("SUPABASE_SERVICE_ROLE_KEY", "")
# Rate limit: 5 requests per second
RATE_LIMIT_MAX = 5
RATE_LIMIT_WINDOW = 1 # seconds
# Input limits
MAX_CODE_LENGTH = 50_000 # 50KB max source code
MAX_STDIN_LENGTH = 10_000 # 10KB max stdin
app = FastAPI(title="Code Execution Gateway", docs_url=None, redoc_url=None)
# CORS — only allow your frontend origin
app.add_middleware(
CORSMiddleware,
allow_origins=[
os.getenv("FRONTEND_URL", "https://your-frontend.vercel.app"),
],
allow_methods=["POST", "GET"],
allow_headers=["Content-Type", "X-API-Key"],
)
redis_conn = Redis(host='localhost', port=6379, decode_responses=True)
# Supabase client
from supabase import create_client, Client
supabase: Client = create_client(SUPABASE_URL, SUPABASE_SERVICE_ROLE_KEY)
# Persistent HTTP client — reuses connections (HTTP keep-alive) for speed
compiler_client = httpx.AsyncClient(
timeout=35.0,
limits=httpx.Limits(max_connections=20, max_keepalive_connections=10),
headers={
"Authorization": COMPILER_API_KEY,
"Content-Type": "application/json",
},
)
# Map our language names → onlinecompiler.io compiler IDs
COMPILER_MAP = {
"python": "python-3.14",
"javascript": "typescript-deno",
"c": "gcc-15",
"cpp": "g++-15",
"java": "openjdk-25",
}
# --- Models ---
class CodeExecutionRequest(BaseModel):
problem_id: str
source_code: str
language: str
is_submit: bool = False
user_id: str
@field_validator("source_code")
@classmethod
def validate_code_length(cls, v):
if len(v) > MAX_CODE_LENGTH:
raise ValueError(f"Source code exceeds {MAX_CODE_LENGTH} character limit")
return v
@field_validator("language")
@classmethod
def validate_language(cls, v):
if v not in COMPILER_MAP:
raise ValueError(f"Unsupported language: {v}. Supported: {list(COMPILER_MAP.keys())}")
return v
class TestCaseResult(BaseModel):
test_case: int
passed: bool
stdout: str
stderr: str
exit_code: int | None = None
timed_out: bool = False
class ExecutionResponse(BaseModel):
job_id: str
status: str
results: list[TestCaseResult] = []
all_passed: bool = False
# --- Rate Limiter ---
def check_rate_limit():
"""Sliding-window rate limiter using Redis INCR + EXPIRE."""
key = f"ratelimit:global:{int(time.time())}"
current = redis_conn.incr(key)
if current == 1:
redis_conn.expire(key, RATE_LIMIT_WINDOW + 1)
if current > RATE_LIMIT_MAX:
raise HTTPException(
status_code=429,
detail="Rate limit exceeded. Try again in a moment."
)
# --- API Key Verification (timing-safe) ---
async def verify_api_key(x_api_key: str = Header(None)):
if not x_api_key or not EXECUTION_API_KEY:
raise HTTPException(status_code=401, detail="Invalid API Key")
if not hmac.compare_digest(x_api_key, EXECUTION_API_KEY):
raise HTTPException(status_code=401, detail="Invalid API Key")
return True
# --- OnlineCompiler.io API Caller ---
async def call_compiler(language: str, source_code: str, stdin: str = "") -> dict:
"""Send code to onlinecompiler.io sync endpoint and return the result."""
compiler = COMPILER_MAP.get(language)
if not compiler:
return {
"stdout": "",
"stderr": f"Unsupported language: {language}",
"exit_code": 1,
"timed_out": False,
}
payload = {
"compiler": compiler,
"code": source_code,
"input": stdin[:MAX_STDIN_LENGTH],
}
resp = await compiler_client.post(COMPILER_API_URL, json=payload)
if resp.status_code != 200:
raise HTTPException(
status_code=502,
detail="Code execution service unavailable. Try again."
)
data = resp.json()
return {
"stdout": data.get("output", ""),
"stderr": data.get("error", ""),
"exit_code": data.get("exit_code", 0),
"timed_out": data.get("signal") == "SIGKILL",
}
# --- Core Endpoint ---
@app.post("/execute", response_model=ExecutionResponse)
async def execute_code(
request: CodeExecutionRequest,
authorized: bool = Depends(verify_api_key),
):
"""
1. Rate-limit check
2. Fetch test cases from Supabase
3. For each test case, call OnlineCompiler API
4. Compare output, write results to Supabase
5. Return results
"""
check_rate_limit()
job_id = f"job_{int(time.time() * 1000)}_{request.user_id[:8]}"
try:
# 1. Fetch test cases
response = supabase.table("problems").select("test_cases").eq("id", request.problem_id).single().execute()
if not response.data:
raise HTTPException(status_code=404, detail=f"Problem {request.problem_id} not found")
test_cases = response.data.get("test_cases", [])
results = []
all_passed = True
# 2. Execute each test case
for index, tc in enumerate(test_cases):
check_rate_limit()
exec_result = await call_compiler(
language=request.language,
source_code=request.source_code,
stdin=tc.get("input", ""),
)
actual_output = exec_result["stdout"].strip()
expected_output = (tc.get("expected") or tc.get("output") or "").strip()
passed = actual_output == expected_output
if not passed:
all_passed = False
results.append(TestCaseResult(
test_case=index + 1,
passed=passed,
stdout=actual_output if not request.is_submit else "Hidden",
stderr=exec_result["stderr"],
exit_code=exec_result["exit_code"],
timed_out=exec_result["timed_out"],
))
# If it's a 'Run' and one fails, stop early
if not passed and not request.is_submit:
break
# 3. Persist to Supabase
final_payload = {
"job_id": job_id,
"user_id": request.user_id,
"problem_id": request.problem_id,
"status": "completed",
"results": [r.model_dump() for r in results],
"all_passed": all_passed,
"created_at": "now()",
}
if request.is_submit:
supabase.table("submissions").insert(final_payload).execute()
supabase.table("job_status").insert({
"job_id": job_id,
"user_id": request.user_id,
"data": final_payload,
}).execute()
return ExecutionResponse(
job_id=job_id,
status="completed",
results=results,
all_passed=all_passed,
)
except HTTPException:
raise
except Exception as e:
supabase.table("job_status").insert({
"job_id": job_id,
"user_id": request.user_id,
"data": {"status": "error", "message": str(e)},
}).execute()
raise HTTPException(status_code=500, detail="Execution failed. Please try again.")
# --- Stress Test Endpoint (no Supabase) ---
class StressTestRequest(BaseModel):
source_code: str = 'print("hello")'
language: str = "python"
stdin: str = ""
@app.post("/stress-test")
async def stress_test(
request: StressTestRequest,
authorized: bool = Depends(verify_api_key),
):
"""Lightweight endpoint for load testing. Skips Supabase entirely."""
check_rate_limit()
result = await call_compiler(
language=request.language,
source_code=request.source_code,
stdin=request.stdin,
)
return {
"status": "ok",
"stdout": result["stdout"].strip(),
"stderr": result["stderr"],
"exit_code": result["exit_code"],
"timed_out": result["timed_out"],
}
@app.get("/health")
async def health_check():
try:
redis_conn.ping()
redis_status = "connected"
except Exception:
redis_status = "disconnected"
return {
"status": "online",
"redis": redis_status,
"compiler_api": COMPILER_API_URL,
}
@app.on_event("shutdown")
async def shutdown():
await compiler_client.aclose()
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=7860)