Spaces:
Sleeping
Sleeping
| import os | |
| from typing import Optional | |
| import gradio as gr | |
| import plotly.graph_objects as go | |
| from authlib.integrations.base_client import OAuthError | |
| from authlib.integrations.starlette_client import OAuth | |
| from dotenv import load_dotenv | |
| from fastapi import FastAPI | |
| from starlette.config import Config | |
| from starlette.middleware.base import BaseHTTPMiddleware | |
| from starlette.middleware.sessions import SessionMiddleware | |
| from starlette.requests import Request | |
| from starlette.responses import HTMLResponse, RedirectResponse | |
| from starlette.staticfiles import StaticFiles | |
| from starlette.templating import Jinja2Templates | |
| load_dotenv() | |
| app = FastAPI(title="Orchestrator Evals OAuth", version="1.0.0") | |
| SECRET_KEY = os.getenv("SECRET_KEY", "change-me") | |
| SPACE_HOST = os.getenv("SPACE_HOST", "").strip() | |
| LOCAL_HTTPS_REDIRECT = os.getenv("LOCAL_HTTPS_REDIRECT", "").strip().lower() in {"1", "true", "yes"} | |
| config = Config(".env") | |
| oauth = OAuth(config) | |
| oauth.register( | |
| name="google", | |
| server_metadata_url="https://accounts.google.com/.well-known/openid-configuration", | |
| client_kwargs={"scope": "openid email profile"}, | |
| ) | |
| def _external_base_url(request: Request) -> str: | |
| # In HF Spaces, SPACE_HOST is the canonical public host. | |
| if SPACE_HOST: | |
| return f"https://{SPACE_HOST}" | |
| base_url = str(request.base_url).rstrip("/") | |
| return base_url | |
| def _build_redirect_uri(request: Request) -> str: | |
| external_base = _external_base_url(request) | |
| redirect_uri = f"{external_base}/auth" | |
| # For local dev with plain uvicorn, keep http unless explicitly opted-in. | |
| if LOCAL_HTTPS_REDIRECT and ("://localhost" in redirect_uri or "://127.0.0.1" in redirect_uri): | |
| redirect_uri = redirect_uri.replace("http://", "https://", 1) | |
| return redirect_uri | |
| def _user_from_session(request: Request) -> Optional[dict]: | |
| if "session" not in request.scope: | |
| return None | |
| user = request.session.get("user") | |
| if isinstance(user, dict): | |
| return user | |
| return None | |
| class GradioAuthMiddleware(BaseHTTPMiddleware): | |
| async def dispatch(self, request: Request, call_next): | |
| if request.url.path.startswith("/gradio"): | |
| if _user_from_session(request) is None: | |
| return RedirectResponse(url="/", status_code=307) | |
| return await call_next(request) | |
| app.add_middleware(GradioAuthMiddleware) | |
| app.add_middleware( | |
| SessionMiddleware, | |
| secret_key=SECRET_KEY, | |
| max_age=3600, | |
| same_site="lax", | |
| https_only=True, | |
| ) | |
| app.mount("/static", StaticFiles(directory="static"), name="static") | |
| templates = Jinja2Templates(directory="templates") | |
| async def homepage(request: Request): | |
| user = _user_from_session(request) | |
| if user is None: | |
| return templates.TemplateResponse( | |
| request=request, | |
| name="home_public.html", | |
| context={"title": "Orchestrator Evals"}, | |
| ) | |
| name = user.get("name") or user.get("email") or "User" | |
| return templates.TemplateResponse( | |
| request=request, | |
| name="home_authenticated.html", | |
| context={"title": "Orchestrator Evals", "name": name}, | |
| ) | |
| async def login(request: Request): | |
| redirect_uri = _build_redirect_uri(request) | |
| return await oauth.google.authorize_redirect(request, redirect_uri) | |
| async def auth(request: Request): | |
| try: | |
| token = await oauth.google.authorize_access_token(request) | |
| user = await oauth.google.userinfo(token=token) | |
| request.session["user"] = dict(user) | |
| return RedirectResponse(url="/gradio", status_code=302) | |
| except OAuthError as exc: | |
| return HTMLResponse(f"OAuth error: {exc}", status_code=400) | |
| except Exception as exc: # pragma: no cover | |
| return HTMLResponse(f"Authentication failed: {exc}", status_code=500) | |
| async def logout(request: Request): | |
| request.session.pop("user", None) | |
| return RedirectResponse(url="/", status_code=302) | |
| async def health(): | |
| return {"status": "ok"} | |
| def _build_dummy_plot() -> go.Figure: | |
| fig = go.Figure( | |
| data=[ | |
| go.Scatter( | |
| x=["Mon", "Tue", "Wed", "Thu", "Fri"], | |
| y=[2, 4, 3, 5, 6], | |
| mode="lines+markers", | |
| name="Demo Series", | |
| ) | |
| ] | |
| ) | |
| fig.update_layout( | |
| title="Dummy Plotly Graph", | |
| xaxis_title="Day", | |
| yaxis_title="Value", | |
| template="plotly_white", | |
| ) | |
| return fig | |
| with gr.Blocks(title="Orchestrator Evals") as gradio_ui: | |
| gr.Markdown("## Protected Gradio UI") | |
| gr.Markdown("This is a dummy Plotly graph behind Google OAuth.") | |
| plot = gr.Plot(label="Plotly Demo") | |
| gradio_ui.load(fn=_build_dummy_plot, outputs=plot) | |
| app = gr.mount_gradio_app(app, gradio_ui, path="/gradio") | |