Spaces:
Running
Running
| """ | |
| AutoGEO Studio — FastAPI server entry point. | |
| Replaces the old Streamlit-based app with a stable HTTP architecture: | |
| FastAPI + Jinja2 + HTMX + Alpine.js + SSE (no WebSocket). | |
| """ | |
| import json | |
| import os | |
| import sys | |
| from pathlib import Path | |
| from fastapi import FastAPI, Request | |
| from fastapi.responses import HTMLResponse, JSONResponse | |
| from fastapi.staticfiles import StaticFiles | |
| from fastapi.templating import Jinja2Templates | |
| # --------------------------------------------------------------------------- | |
| # Ensure parent AND app dir are on path | |
| # --------------------------------------------------------------------------- | |
| APP_DIR = Path(__file__).resolve().parent | |
| ROOT_DIR = APP_DIR.parent | |
| for p in (str(APP_DIR), str(ROOT_DIR)): | |
| if p not in sys.path: | |
| sys.path.insert(0, p) | |
| # --------------------------------------------------------------------------- | |
| # Load API keys BEFORE importing engine/evaluation modules | |
| # (they initialise API clients at import time) | |
| # --------------------------------------------------------------------------- | |
| from config_ui import load_keys, save_keys, apply_keys_to_env # noqa: E402 | |
| _keys = load_keys() | |
| for k, v in _keys.items(): | |
| if v: | |
| os.environ[k] = v | |
| from db import init_db # noqa: E402 | |
| # --------------------------------------------------------------------------- | |
| # FastAPI app | |
| # --------------------------------------------------------------------------- | |
| app = FastAPI(title="AutoGEO Studio", version="2.0.0") | |
| # Static files (CSS / JS) | |
| static_dir = APP_DIR / "static" | |
| static_dir.mkdir(parents=True, exist_ok=True) | |
| app.mount("/static", StaticFiles(directory=str(static_dir)), name="static") | |
| # Jinja2 templates | |
| templates = Jinja2Templates(directory=str(APP_DIR / "templates")) | |
| # --------------------------------------------------------------------------- | |
| # Startup | |
| # --------------------------------------------------------------------------- | |
| init_db() | |
| # --------------------------------------------------------------------------- | |
| # Helpers | |
| # --------------------------------------------------------------------------- | |
| def _is_htmx(request: Request) -> bool: | |
| """True when the request comes from HTMX (fragment swap).""" | |
| return request.headers.get("HX-Request", "").lower() == "true" | |
| def _render_page(request: Request, partial: str, context: dict | None = None) -> HTMLResponse: | |
| """Return full page (base.html) when first load, or fragment when HTMX.""" | |
| ctx = context or {} | |
| # Inject sidebar context | |
| ctx["keys"] = load_keys() | |
| ctx["engine_llm"] = _get_active_engine(request) | |
| if _is_htmx(request): | |
| return templates.TemplateResponse(request=request, name=f"pages/{partial}.html", context=ctx) | |
| # Pass the active template path so base.html can include it dynamically | |
| ctx["active_page"] = partial | |
| ctx["active_template"] = f"pages/{partial}.html" | |
| return templates.TemplateResponse(request=request, name="base.html", context=ctx) | |
| def _get_active_engine(request: Request) -> str: | |
| """Extract engine_llm from cookies, fallback to doubao.""" | |
| engine = request.cookies.get("engine_llm", "doubao") | |
| return engine | |
| # --------------------------------------------------------------------------- | |
| # Import and register API routers | |
| # --------------------------------------------------------------------------- | |
| from api.history import router as history_router # noqa: E402 | |
| from api.distractors import router as distractor_router # noqa: E402 | |
| from api.optimizer import router as optimizer_router # noqa: E402 | |
| app.include_router(history_router) | |
| app.include_router(distractor_router) | |
| app.include_router(optimizer_router) | |
| # ============================================================================= | |
| # PAGE ROUTES (HTML) | |
| # ============================================================================= | |
| async def page_home(request: Request): | |
| return _render_page(request, "home") | |
| async def page_distractors(request: Request): | |
| return _render_page(request, "distractors") | |
| async def page_optimize(request: Request): | |
| from db import get_all_distractor_sets | |
| sets = get_all_distractor_sets() | |
| ctx = {"saved_sets": sets} | |
| return _render_page(request, "optimize", ctx) | |
| async def page_history(request: Request): | |
| from db import get_all_records | |
| records = get_all_records() | |
| ctx = {"records": records} | |
| return _render_page(request, "history", ctx) | |
| # ============================================================================= | |
| # API KEY ENDPOINT (form-based, no separate router) | |
| # ============================================================================= | |
| async def update_keys(request: Request): | |
| """Handle sidebar API-key form submission. Returns updated sidebar.""" | |
| form = await request.form() | |
| new_keys = { | |
| "DEEPSEEK_API_KEY": form.get("DEEPSEEK_API_KEY", ""), | |
| "DOUBAO_API_KEY": form.get("DOUBAO_API_KEY", ""), | |
| "DOUBAO_ENDPOINT_ID": form.get("DOUBAO_ENDPOINT_ID", "ep-20260525112247-ccwhk"), | |
| "OPENAI_API_KEY": form.get("OPENAI_API_KEY", ""), | |
| "GEMINI_API_KEY": form.get("GEMINI_API_KEY", ""), | |
| "ANTHROPIC_API_KEY": form.get("ANTHROPIC_API_KEY", ""), | |
| "HTTP_PROXY": form.get("HTTP_PROXY", ""), | |
| "HTTPS_PROXY": form.get("HTTPS_PROXY", ""), | |
| } | |
| save_keys(new_keys) | |
| apply_keys_to_env(new_keys) | |
| # Set cookie for engine selection | |
| engine_llm = form.get("engine_llm", "doubao") | |
| response = templates.TemplateResponse( | |
| request=request, name="components/keys_saved.html", context={"keys": new_keys}, | |
| ) | |
| response.set_cookie(key="engine_llm", value=engine_llm) | |
| return response | |
| # ============================================================================= | |
| # MAIN | |
| # ============================================================================= | |
| if __name__ == "__main__": | |
| import uvicorn | |
| host = os.getenv("HOST", "0.0.0.0") | |
| port = int(os.getenv("PORT", "8000")) | |
| uvicorn.run("server:app", host=host, port=port, reload=True) | |