From 741b7d7b4a21764a9884a013acfe42c7ad66d49d Mon Sep 17 00:00:00 2001 From: Tal Ben-Nun Date: Fri, 5 Dec 2025 15:11:22 -0800 Subject: [PATCH] Refactor server into routers and add database backend --- .gitignore | 1 + backend/__init__.py | 0 backend/auth.py | 19 ++ backend/database/__init__.py | 0 backend/database/engine.py | 51 ++++++ backend/database/models.py | 71 ++++++++ backend/database/schemas.py | 109 +++++++++++ backend/routers/__init__.py | 0 backend/routers/projects.py | 243 +++++++++++++++++++++++++ backend/routers/webui.py | 45 +++++ backend/server.py | 50 +++++ charge_backend/backend_helper_funcs.py | 58 +----- charge_backend/charge_server.py | 97 +--------- charge_backend/structures.py | 59 ++++++ dockerscripts/launch_servers.sh | 2 +- mock_server.py | 98 +--------- requirements.txt | 3 + 17 files changed, 673 insertions(+), 233 deletions(-) create mode 100644 backend/__init__.py create mode 100644 backend/auth.py create mode 100644 backend/database/__init__.py create mode 100644 backend/database/engine.py create mode 100644 backend/database/models.py create mode 100644 backend/database/schemas.py create mode 100644 backend/routers/__init__.py create mode 100644 backend/routers/projects.py create mode 100644 backend/routers/webui.py create mode 100644 backend/server.py create mode 100644 charge_backend/structures.py diff --git a/.gitignore b/.gitignore index 08972f38..04ce48ac 100644 --- a/.gitignore +++ b/.gitignore @@ -142,6 +142,7 @@ vite.config.ts.timestamp-* *~ .venv/ venv/ +__pycache__/ # Visual Studio Code .vscode/ diff --git a/backend/__init__.py b/backend/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/backend/auth.py b/backend/auth.py new file mode 100644 index 00000000..de16eeaa --- /dev/null +++ b/backend/auth.py @@ -0,0 +1,19 @@ +################################################################################ +## Copyright 2025 Lawrence Livermore National Security, LLC. and Binghamton University. +## See the top-level LICENSE file for details. +## +## SPDX-License-Identifier: Apache-2.0 +################################################################################ + +from fastapi import Header, HTTPException, status +from typing import Optional + + +async def get_forwarded_user(x_forwarded_user: Optional[str] = Header(None, alias="X-Forwarded-User")) -> str: + """Extract authenticated user from X-Forwarded-User header""" + if not x_forwarded_user: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Missing X-Forwarded-User header. Authentication required.", + ) + return x_forwarded_user diff --git a/backend/database/__init__.py b/backend/database/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/backend/database/engine.py b/backend/database/engine.py new file mode 100644 index 00000000..194a65ae --- /dev/null +++ b/backend/database/engine.py @@ -0,0 +1,51 @@ +################################################################################ +## Copyright 2025 Lawrence Livermore National Security, LLC. and Binghamton University. +## See the top-level LICENSE file for details. +## +## SPDX-License-Identifier: Apache-2.0 +########################################################################## + +from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker +from sqlalchemy.exc import OperationalError +from sqlalchemy.orm import declarative_base +import os + +if "MARIADB_HOST" not in os.environ: + engine = AsyncSessionLocal = None +else: + DB_USER = os.getenv("MARIADB_USER", "user") + DB_PASSWORD = os.getenv("MARIADB_PASSWORD", "password") + DB_HOST = os.getenv("MARIADB_HOST", "localhost") + DB_PORT = os.getenv("MARIADB_PORT", "8080") + DATABASE_URL = f"mysql+aiomysql://{DB_USER}:{DB_PASSWORD}@{DB_HOST}:{DB_PORT}" + + try: + engine = create_async_engine( + DATABASE_URL, + echo=True, + pool_size=10, + max_overflow=20, + pool_pre_ping=True, + pool_recycle=3600, + ) + + AsyncSessionLocal = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) + except OperationalError: + engine = AsyncSessionLocal = None + +Base = declarative_base() + + +async def get_db(): + if AsyncSessionLocal is None: + yield None + return + async with AsyncSessionLocal() as session: + try: + yield session + await session.commit() + except Exception: + await session.rollback() + raise + finally: + await session.close() diff --git a/backend/database/models.py b/backend/database/models.py new file mode 100644 index 00000000..b0bc610b --- /dev/null +++ b/backend/database/models.py @@ -0,0 +1,71 @@ +################################################################################ +## Copyright 2025 Lawrence Livermore National Security, LLC. and Binghamton University. +## See the top-level LICENSE file for details. +## +## SPDX-License-Identifier: Apache-2.0 +################################################################################ + +from sqlalchemy import Column, String, DateTime, Boolean, ForeignKey, Index, Text, JSON, Float +from sqlalchemy.orm import relationship, Mapped, mapped_column +from datetime import datetime +from typing import Optional +from backend.database.engine import Base + + +class Project(Base): + __tablename__ = "projects" + + id: Mapped[str] = mapped_column(String(255), primary_key=True) + user: Mapped[str] = mapped_column(String(255), nullable=False, index=True) + name: Mapped[str] = mapped_column(String(255), nullable=False) + created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, nullable=False) + last_modified: Mapped[datetime] = mapped_column( + DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False + ) + + experiments: Mapped[list["Experiment"]] = relationship( + "Experiment", back_populates="project", cascade="all, delete-orphan" + ) + + __table_args__ = (Index("idx_user_last_modified", "user", "last_modified"),) + + +class Experiment(Base): + __tablename__ = "experiments" + + id: Mapped[str] = mapped_column(String(255), primary_key=True) + project_id: Mapped[str] = mapped_column(String(255), ForeignKey("projects.id", ondelete="CASCADE"), nullable=False) + user: Mapped[str] = mapped_column(String(255), nullable=False, index=True) + name: Mapped[str] = mapped_column(String(255), nullable=False) + created_at: Mapped[datetime] = mapped_column(DateTime, default=datetime.utcnow, nullable=False) + last_modified: Mapped[datetime] = mapped_column( + DateTime, default=datetime.utcnow, onupdate=datetime.utcnow, nullable=False + ) + is_running: Mapped[Optional[bool]] = mapped_column(Boolean, default=False) + + # System state fields + smiles: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + problem_type: Mapped[Optional[str]] = mapped_column(String(255), nullable=True) + problem_name: Mapped[Optional[str]] = mapped_column(String(255), nullable=True) + system_prompt: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + problem_prompt: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + property_type: Mapped[Optional[str]] = mapped_column(String(255), nullable=True) + custom_property_name: Mapped[Optional[str]] = mapped_column(String(255), nullable=True) + custom_property_desc: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + custom_property_ascending: Mapped[Optional[bool]] = mapped_column(Boolean, nullable=True) + + # Complex JSON fields for nested data structures + tree_nodes: Mapped[Optional[str]] = mapped_column(JSON, nullable=True) + edges: Mapped[Optional[str]] = mapped_column(JSON, nullable=True) + metrics_history: Mapped[Optional[str]] = mapped_column(JSON, nullable=True) + visible_metrics: Mapped[Optional[str]] = mapped_column(JSON, nullable=True) + graph_state: Mapped[Optional[str]] = mapped_column(JSON, nullable=True) + auto_zoom: Mapped[Optional[bool]] = mapped_column(Boolean, nullable=True) + sidebar_state: Mapped[Optional[str]] = mapped_column(JSON, nullable=True) + + # Experiment context + experiment_context: Mapped[Optional[str]] = mapped_column(Text, nullable=True) + + project: Mapped["Project"] = relationship("Project", back_populates="experiments") + + __table_args__ = (Index("idx_user_project", "user", "project_id"),) diff --git a/backend/database/schemas.py b/backend/database/schemas.py new file mode 100644 index 00000000..71a63432 --- /dev/null +++ b/backend/database/schemas.py @@ -0,0 +1,109 @@ +################################################################################ +## Copyright 2025 Lawrence Livermore National Security, LLC. and Binghamton University. +## See the top-level LICENSE file for details. +## +## SPDX-License-Identifier: Apache-2.0 +################################################################################ + +from pydantic import BaseModel, Field +from datetime import datetime +from typing import List, Optional, Any + + +class ExperimentBase(BaseModel): + name: str + + +class ExperimentCreate(ExperimentBase): + pass + + +class ExperimentUpdate(BaseModel): + name: Optional[str] = None + is_running: Optional[bool] = None + + # System state fields + smiles: Optional[str] = None + problem_type: Optional[str] = Field(None, alias="problemType") + problem_name: Optional[str] = Field(None, alias="problemName") + system_prompt: Optional[str] = Field(None, alias="systemPrompt") + problem_prompt: Optional[str] = Field(None, alias="problemPrompt") + property_type: Optional[str] = Field(None, alias="propertyType") + custom_property_name: Optional[str] = Field(None, alias="customPropertyName") + custom_property_desc: Optional[str] = Field(None, alias="customPropertyDesc") + custom_property_ascending: Optional[bool] = Field(None, alias="customPropertyAscending") + + # Complex nested data (stored as JSON) + tree_nodes: Optional[Any] = Field(None, alias="treeNodes") + edges: Optional[Any] = None + metrics_history: Optional[Any] = Field(None, alias="metricsHistory") + visible_metrics: Optional[Any] = Field(None, alias="visibleMetrics") + graph_state: Optional[Any] = Field(None, alias="graphState") + auto_zoom: Optional[bool] = Field(None, alias="autoZoom") + sidebar_state: Optional[Any] = Field(None, alias="sidebarState") + + # Experiment context + experiment_context: Optional[str] = Field(None, alias="experimentContext") + + class Config: + populate_by_name = True + + +class Experiment(ExperimentBase): + id: str + project_id: str = Field(..., alias="projectId") + user: str + created_at: datetime = Field(..., alias="createdAt") + last_modified: datetime = Field(..., alias="lastModified") + is_running: Optional[bool] = Field(None, alias="isRunning") + + # System state fields + smiles: Optional[str] = None + problem_type: Optional[str] = Field(None, alias="problemType") + problem_name: Optional[str] = Field(None, alias="problemName") + system_prompt: Optional[str] = Field(None, alias="systemPrompt") + problem_prompt: Optional[str] = Field(None, alias="problemPrompt") + property_type: Optional[str] = Field(None, alias="propertyType") + custom_property_name: Optional[str] = Field(None, alias="customPropertyName") + custom_property_desc: Optional[str] = Field(None, alias="customPropertyDesc") + custom_property_ascending: Optional[bool] = Field(None, alias="customPropertyAscending") + + # Complex nested data + tree_nodes: Optional[Any] = Field(None, alias="treeNodes") + edges: Optional[Any] = None + metrics_history: Optional[Any] = Field(None, alias="metricsHistory") + visible_metrics: Optional[Any] = Field(None, alias="visibleMetrics") + graph_state: Optional[Any] = Field(None, alias="graphState") + auto_zoom: Optional[bool] = Field(None, alias="autoZoom") + sidebar_state: Optional[Any] = Field(None, alias="sidebarState") + + # Experiment context + experiment_context: Optional[str] = Field(None, alias="experimentContext") + + class Config: + from_attributes = True + populate_by_name = True + + +class ProjectBase(BaseModel): + name: str + + +class ProjectCreate(ProjectBase): + pass + + +class ProjectUpdate(ProjectBase): + pass + + +class Project(ProjectBase): + id: str + user: str + created_at: datetime = Field(..., alias="createdAt") + last_modified: datetime = Field(..., alias="lastModified") + experiments: List[Experiment] = [] + + class Config: + from_attributes = True + populate_by_name = True diff --git a/backend/routers/__init__.py b/backend/routers/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/backend/routers/projects.py b/backend/routers/projects.py new file mode 100644 index 00000000..02cbf06a --- /dev/null +++ b/backend/routers/projects.py @@ -0,0 +1,243 @@ +################################################################################ +## Copyright 2025 Lawrence Livermore National Security, LLC. and Binghamton University. +## See the top-level LICENSE file for details. +## +## SPDX-License-Identifier: Apache-2.0 +################################################################################ +""" +Serves routes for database interaction for projects and experiments (``/api/projects/*``). +""" +from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy import select +from sqlalchemy.orm import selectinload +from typing import List +import time +import random +from datetime import datetime + +from backend.database.engine import get_db +from backend.auth import get_forwarded_user +from backend.database import models, schemas + +router = APIRouter(prefix="/api/projects", tags=["projects"]) + + +def generate_id(prefix: str) -> str: + """Generate unique ID""" + timestamp = int(time.time() * 1000) + random_str = "".join(random.choices("abcdefghijklmnopqrstuvwxyz0123456789", k=9)) + return f"{prefix}_{timestamp}_{random_str}" + + +@router.get("/", response_model=List[schemas.Project]) +async def get_projects(user: str = Depends(get_forwarded_user), db: AsyncSession = Depends(get_db)): + """Load all projects for the authenticated user""" + result = await db.execute( + select(models.Project) + .where(models.Project.user == user) + .options(selectinload(models.Project.experiments)) + .order_by(models.Project.last_modified.desc()) + ) + projects = result.scalars().all() + return projects + + +@router.post("/", response_model=schemas.Project) +async def create_project( + project_data: schemas.ProjectCreate, user: str = Depends(get_forwarded_user), db: AsyncSession = Depends(get_db) +): + """Create a new project for the authenticated user""" + new_project = models.Project( + id=generate_id("project"), + user=user, + name=project_data.name, + created_at=datetime.utcnow(), + last_modified=datetime.utcnow(), + ) + db.add(new_project) + await db.commit() + await db.refresh(new_project) + return new_project + + +@router.put("/{project_id}", response_model=schemas.Project) +async def update_project( + project_id: str, + project_data: schemas.ProjectUpdate, + user: str = Depends(get_forwarded_user), + db: AsyncSession = Depends(get_db), +): + """Update a project (only if owned by user)""" + result = await db.execute( + select(models.Project).where(models.Project.id == project_id, models.Project.user == user) + ) + project = result.scalar_one_or_none() + + if not project: + raise HTTPException(status_code=404, detail="Project not found") + + project.name = project_data.name + project.last_modified = datetime.utcnow() + + await db.commit() + await db.refresh(project) + return project + + +@router.delete("/{project_id}") +async def delete_project(project_id: str, user: str = Depends(get_forwarded_user), db: AsyncSession = Depends(get_db)): + """Delete a project (only if owned by user)""" + result = await db.execute( + select(models.Project).where(models.Project.id == project_id, models.Project.user == user) + ) + project = result.scalar_one_or_none() + + if not project: + raise HTTPException(status_code=404, detail="Project not found") + + await db.delete(project) + await db.commit() + return {"success": True} + + +@router.post("/{project_id}/experiments", response_model=schemas.Experiment) +async def create_experiment( + project_id: str, + experiment_data: schemas.ExperimentCreate, + user: str = Depends(get_forwarded_user), + db: AsyncSession = Depends(get_db), +): + """Create a new experiment (only if user owns project)""" + result = await db.execute( + select(models.Project).where(models.Project.id == project_id, models.Project.user == user) + ) + project = result.scalar_one_or_none() + + if not project: + raise HTTPException(status_code=404, detail="Project not found") + + new_experiment = models.Experiment( + id=generate_id("exp"), + project_id=project_id, + user=user, + name=experiment_data.name, + created_at=datetime.utcnow(), + last_modified=datetime.utcnow(), + ) + + project.last_modified = datetime.utcnow() + + db.add(new_experiment) + await db.commit() + await db.refresh(new_experiment) + return new_experiment + + +@router.put("/{project_id}/experiments/{experiment_id}", response_model=schemas.Experiment) +async def update_experiment( + project_id: str, + experiment_id: str, + experiment_data: schemas.ExperimentUpdate, + user: str = Depends(get_forwarded_user), + db: AsyncSession = Depends(get_db), +): + """Update an experiment (only if user owns it)""" + result = await db.execute( + select(models.Experiment).where( + models.Experiment.id == experiment_id, + models.Experiment.project_id == project_id, + models.Experiment.user == user, + ) + ) + experiment = result.scalar_one_or_none() + + if not experiment: + raise HTTPException(status_code=404, detail="Experiment not found") + + # Update all provided fields + update_data = experiment_data.model_dump(exclude_unset=True, by_alias=False) + for field, value in update_data.items(): + if hasattr(experiment, field): + setattr(experiment, field, value) + + experiment.last_modified = datetime.utcnow() + + # Update project's last_modified + result = await db.execute( + select(models.Project).where(models.Project.id == project_id, models.Project.user == user) + ) + project = result.scalar_one_or_none() + if project: + project.last_modified = datetime.utcnow() + + await db.commit() + await db.refresh(experiment) + return experiment + + +@router.delete("/{project_id}/experiments/{experiment_id}") +async def delete_experiment( + project_id: str, experiment_id: str, user: str = Depends(get_forwarded_user), db: AsyncSession = Depends(get_db) +): + """Delete an experiment (only if user owns it)""" + result = await db.execute( + select(models.Experiment).where( + models.Experiment.id == experiment_id, + models.Experiment.project_id == project_id, + models.Experiment.user == user, + ) + ) + experiment = result.scalar_one_or_none() + + if not experiment: + raise HTTPException(status_code=404, detail="Experiment not found") + + await db.delete(experiment) + + # Update project's last_modified + result = await db.execute( + select(models.Project).where(models.Project.id == project_id, models.Project.user == user) + ) + project = result.scalar_one_or_none() + if project: + project.last_modified = datetime.utcnow() + + await db.commit() + return {"success": True} + + +@router.patch("/{project_id}/experiments/{experiment_id}/running") +async def set_experiment_running( + project_id: str, + experiment_id: str, + is_running: bool, + user: str = Depends(get_forwarded_user), + db: AsyncSession = Depends(get_db), +): + """Set experiment running status (only if user owns it)""" + result = await db.execute( + select(models.Experiment).where( + models.Experiment.id == experiment_id, + models.Experiment.project_id == project_id, + models.Experiment.user == user, + ) + ) + experiment = result.scalar_one_or_none() + + if not experiment: + raise HTTPException(status_code=404, detail="Experiment not found") + + experiment.is_running = is_running + experiment.last_modified = datetime.utcnow() + + # Update project's last_modified + result = await db.execute( + select(models.Project).where(models.Project.id == project_id, models.Project.user == user) + ) + project = result.scalar_one_or_none() + if project: + project.last_modified = datetime.utcnow() + + await db.commit() + return {"success": True, "is_running": is_running} diff --git a/backend/routers/webui.py b/backend/routers/webui.py new file mode 100644 index 00000000..cc7b6717 --- /dev/null +++ b/backend/routers/webui.py @@ -0,0 +1,45 @@ +################################################################################ +## Copyright 2025 Lawrence Livermore National Security, LLC. and Binghamton University. +## See the top-level LICENSE file for details. +## +## SPDX-License-Identifier: Apache-2.0 +################################################################################ +""" +Web UI (``/``) routing +""" +from fastapi import Request, APIRouter +from fastapi.responses import HTMLResponse +from fastapi.staticfiles import StaticFiles +import os +from loguru import logger + +router = APIRouter(tags=["webui"]) + +if "FLASK_APPDIR" in os.environ: + DIST_PATH = os.environ["FLASK_APPDIR"] +else: + DIST_PATH = os.path.join(os.path.dirname(__file__), "flask-app", "dist") +ASSETS_PATH = os.path.join(DIST_PATH, "assets") + +if os.path.exists(ASSETS_PATH): + # Serve the frontend + router.mount("/assets", StaticFiles(directory=ASSETS_PATH), name="assets") + router.mount("/rdkit", StaticFiles(directory=os.path.join(DIST_PATH, "rdkit")), name="rdkit") + + @router.get("/") + async def root(request: Request): + logger.info(f"Request for Web UI received. Headers: {str(request.headers)}") + with open(os.path.join(DIST_PATH, "index.html"), "r") as fp: + html = fp.read() + + html = html.replace( + "", + f""" + """, + ) + return HTMLResponse(html) diff --git a/backend/server.py b/backend/server.py new file mode 100644 index 00000000..fd612b8a --- /dev/null +++ b/backend/server.py @@ -0,0 +1,50 @@ +################################################################################ +## Copyright 2025 Lawrence Livermore National Security, LLC. and Binghamton University. +## See the top-level LICENSE file for details. +## +## SPDX-License-Identifier: Apache-2.0 +################################################################################ +""" +General server and ``/`` routing infrastructure. +""" +from fastapi import FastAPI +from fastapi.middleware.cors import CORSMiddleware +from contextlib import asynccontextmanager + +from backend.database.engine import engine, Base +from backend.routers import webui + +@asynccontextmanager +async def lifespan(app: FastAPI): + if engine is not None: + # Startup: Create tables + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + + yield + + if engine is not None: + # Shutdown: Close connections + await engine.dispose() + + +app = FastAPI(title="Flask Copilot Backend", lifespan=lifespan) + +app.add_middleware( + CORSMiddleware, + allow_origins=["*"], + allow_credentials=True, + allow_methods=["*"], + allow_headers=["*"], +) + +@app.get("/health") +async def health_check(): + return {"status": "healthy"} + +app.include_router(webui.router) + +if __name__ == "__main__": + import uvicorn + + uvicorn.run(app, host="0.0.0.0", port=8001) diff --git a/charge_backend/backend_helper_funcs.py b/charge_backend/backend_helper_funcs.py index ef81488b..d5dc072b 100644 --- a/charge_backend/backend_helper_funcs.py +++ b/charge_backend/backend_helper_funcs.py @@ -1,9 +1,14 @@ +################################################################################ +## Copyright 2025 Lawrence Livermore National Security, LLC. and Binghamton University. +## See the top-level LICENSE file for details. +## +## SPDX-License-Identifier: Apache-2.0 +################################################################################ from loguru import logger from fastapi import WebSocket import asyncio import json from typing import Dict, Optional, Literal, Tuple -from dataclasses import dataclass, asdict from collections import defaultdict from charge.clients.autogen import AutoGenAgent @@ -11,56 +16,7 @@ from charge.servers import SMILES_utils from charge.servers.molecular_property_utils import get_density from callback_logger import CallbackLogger - - -@dataclass -class Node: - id: str - smiles: str - label: str - hoverInfo: str - level: int - parentId: Optional[str] = None - x: Optional[int] = None - y: Optional[int] = None - # Properties - cost: Optional[float] = None - bandgap: Optional[float] = None - yield_: Optional[float] = None - highlight: Optional[str] = "normal" - density: Optional[float] = None - sascore: Optional[float] = None - - def json(self): - ret = asdict(self) - ret["yield"] = ret["yield_"] - del ret["yield_"] - return ret - - -@dataclass -class Edge: - id: str - fromNode: str - toNode: str - status: Literal["computing", "complete"] - label: Optional[str] = None - - def json(self): - return asdict(self) - - -@dataclass -class ModelMessage: - message: str - smiles: Optional[str] - - def json(self): - ret = asdict(self) - if self.smiles is None: - if "smiles" in ret: - del ret["smiles"] - return ret +from structures import Node, Edge, ModelMessage def get_price(smiles: str) -> float: diff --git a/charge_backend/charge_server.py b/charge_backend/charge_server.py index e3d7b198..b9952881 100644 --- a/charge_backend/charge_server.py +++ b/charge_backend/charge_server.py @@ -5,24 +5,15 @@ ## SPDX-License-Identifier: Apache-2.0 ################################################################################ -import asyncio -from concurrent.futures import ProcessPoolExecutor from functools import partial -from fastapi import FastAPI, WebSocket, WebSocketDisconnect, Request -from fastapi.responses import HTMLResponse -from fastapi.staticfiles import StaticFiles -from fastapi.middleware.cors import CORSMiddleware +from fastapi import WebSocket, WebSocketDisconnect import os import argparse import httpx -from charge.servers.server_utils import try_get_public_hostname import os - - -import logging -from aizynthfinder.utils.logging import setup_logger - from loguru import logger + +from charge.servers.server_utils import try_get_public_hostname from charge.clients.Client import Client from charge.experiments.AutoGenExperiment import AutoGenExperiment from charge.clients.autogen import AutoGenPool @@ -36,6 +27,10 @@ ) from backend_manager import TaskManager, ActionManager +from backend.server import app +from backend.routers import projects as projrouter + +app.include_router(projrouter.router) parser = argparse.ArgumentParser() @@ -69,29 +64,10 @@ ) # Add standard CLI arguments -Client.add_std_parser_arguments( - parser, defaults=dict(backend="openai", model="gpt-5-nano") -) +Client.add_std_parser_arguments(parser, defaults=dict(backend="openai", model="gpt-5-nano")) args, _ = parser.parse_known_args() -app = FastAPI() - -# CORS for development -app.add_middleware( - CORSMiddleware, - allow_origins=["*"], - allow_credentials=True, - allow_methods=["*"], - allow_headers=["*"], -) - -if "FLASK_APPDIR" in os.environ: - DIST_PATH = os.environ["FLASK_APPDIR"] -else: - DIST_PATH = os.path.join(os.path.dirname(__file__), "flask-app", "dist") -ASSETS_PATH = os.path.join(DIST_PATH, "assets") - reload_server_list(args.tool_server_cache) app.post("/register")(partial(register_post, args.tool_server_cache)) @@ -108,55 +84,6 @@ logger.info(f"{status}") count += 1 -if os.path.exists(ASSETS_PATH): - # Serve the frontend - app.mount("/assets", StaticFiles(directory=ASSETS_PATH), name="assets") - app.mount( - "/rdkit", StaticFiles(directory=os.path.join(DIST_PATH, "rdkit")), name="rdkit" - ) - - @app.get("/") - async def root(request: Request): - logger.info(f"Request for Web UI received. Headers: {str(request.headers)}") - with open(os.path.join(DIST_PATH, "index.html"), "r") as fp: - html = fp.read() - - html = html.replace( - "", - f""" - """, - ) - return HTMLResponse(html) - - -async def _cancel_task_if_running( - action, task: asyncio.Task | None, executor: ProcessPoolExecutor | None -): - if action not in [ - "compute", - "optimize-from", - "compute-reaction-from", - "recompute-reaction", - "recompute-parent-reaction", - ]: - return - - if task and not task.done(): - logger.info("Cancelling existing compute task...") - task.cancel() - try: - await task - except asyncio.CancelledError: - logger.info("Previous compute task cancelled.") - - if executor: - executor.shutdown(wait=False, cancel_futures=True) - @app.websocket("/ws") async def websocket_endpoint(websocket: WebSocket): @@ -180,9 +107,7 @@ async def websocket_endpoint(websocket: WebSocket): # set up an AutoGenAgent pool for tasks on this endpoint - autogen_pool = AutoGenPool( - model=model, backend=backend, api_key=API_KEY, base_url=BASE_URL - ) + autogen_pool = AutoGenPool(model=model, backend=backend, api_key=API_KEY, base_url=BASE_URL) # Set up an experiment class for current endpoint experiment = AutoGenExperiment(task=None, agent_pool=autogen_pool) @@ -223,9 +148,7 @@ async def websocket_endpoint(websocket: WebSocket): handler_func = action_handlers[action] await handler_func(data) else: - logger.warning( - f"Unknown action received: {action} with data {data}" - ) + logger.warning(f"Unknown action received: {action} with data {data}") except ValueError as e: logger.error(f"Error in internal loop connection: {e}") await task_manager.cancel_current_task() diff --git a/charge_backend/structures.py b/charge_backend/structures.py new file mode 100644 index 00000000..66041817 --- /dev/null +++ b/charge_backend/structures.py @@ -0,0 +1,59 @@ +################################################################################ +## Copyright 2025 Lawrence Livermore National Security, LLC. and Binghamton University. +## See the top-level LICENSE file for details. +## +## SPDX-License-Identifier: Apache-2.0 +################################################################################ + +from dataclasses import dataclass, asdict +from typing import Optional, Literal + + +@dataclass +class Node: + id: str + smiles: str + label: str + hoverInfo: str + level: int + parentId: Optional[str] = None + x: Optional[int] = None + y: Optional[int] = None + # Properties + cost: Optional[float] = None + bandgap: Optional[float] = None + yield_: Optional[float] = None + highlight: Optional[str] = "normal" + density: Optional[float] = None + sascore: Optional[float] = None + + def json(self): + ret = asdict(self) + ret["yield"] = ret["yield_"] + del ret["yield_"] + return ret + + +@dataclass +class Edge: + id: str + fromNode: str + toNode: str + status: Literal["computing", "complete"] + label: Optional[str] = None + + def json(self): + return asdict(self) + + +@dataclass +class ModelMessage: + message: str + smiles: Optional[str] + + def json(self): + ret = asdict(self) + if self.smiles is None: + if "smiles" in ret: + del ret["smiles"] + return ret diff --git a/dockerscripts/launch_servers.sh b/dockerscripts/launch_servers.sh index 16c0dceb..6d60ad93 100644 --- a/dockerscripts/launch_servers.sh +++ b/dockerscripts/launch_servers.sh @@ -3,7 +3,7 @@ . /venv/bin/activate # Launch main server -export PYTHONPATH=$PYTHONPATH:/app/charge_backend +export PYTHONPATH=$PYTHONPATH:/app/charge_backend:/app uvicorn --host 0.0.0.0 --port 8001 --workers 8 charge_server:app & # Wait for server to start up diff --git a/mock_server.py b/mock_server.py index 7180f12f..a98d0b23 100644 --- a/mock_server.py +++ b/mock_server.py @@ -15,107 +15,17 @@ * ``custom_query``: Execute custom user query """ -from fastapi import FastAPI, Request, WebSocket, WebSocketDisconnect -from fastapi.responses import JSONResponse, HTMLResponse -from fastapi.staticfiles import StaticFiles -from fastapi.middleware.cors import CORSMiddleware +from fastapi import WebSocket, WebSocketDisconnect from dataclasses import dataclass, asdict import asyncio import copy -import os import random -from typing import Any, Optional, Literal -import requests +from typing import Optional from datetime import datetime from loguru import logger from charge_backend.molecule_naming import smiles_to_html - -app = FastAPI() - -# CORS for development -app.add_middleware( - CORSMiddleware, - allow_origins=["*"], - allow_credentials=True, - allow_methods=["*"], - allow_headers=["*"], -) - -if "FLASK_APPDIR" in os.environ: - DIST_PATH = os.environ["FLASK_APPDIR"] -else: - DIST_PATH = os.path.join(os.path.dirname(__file__), "flask-app", "dist") -ASSETS_PATH = os.path.join(DIST_PATH, "assets") - -if os.path.exists(ASSETS_PATH): - # Serve the frontend - app.mount("/assets", StaticFiles(directory=ASSETS_PATH), name="assets") - app.mount( - "/rdkit", StaticFiles(directory=os.path.join(DIST_PATH, "rdkit")), name="rdkit" - ) - - @app.get("/") - async def root(request: Request): - logger.info(f"Request for Web UI received. Headers: {str(request.headers)}") - with open(os.path.join(DIST_PATH, "index.html"), "r") as fp: - html = fp.read() - - html = html.replace( - "", - f""" - """, - ) - return HTMLResponse(html) - - -@dataclass -class Node: - id: str - smiles: str - label: str - hoverInfo: str - level: int - parentId: Optional[str] = None - x: Optional[int] = None - y: Optional[int] = None - # Properties - cost: Optional[float] = None - bandgap: Optional[float] = None - density: Optional[float] = None - yield_: Optional[float] = None - highlight: Optional[str] = "normal" - - def json(self): - ret = asdict(self) - ret["yield"] = ret["yield_"] - del ret["yield_"] - return ret - - -@dataclass -class Edge: - id: str - fromNode: str - toNode: str - status: Literal["computing", "complete"] - label: Optional[str] = None - - def json(self): - return asdict(self) - - -@dataclass -class SidebarMessage: - content: str - smiles: str - - def json(self): - return asdict(self) +from charge_backend.structures import Node, Edge +from backend.server import app @dataclass diff --git a/requirements.txt b/requirements.txt index d38b8a3c..9091331e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,3 +3,6 @@ uvicorn python-multipart websockets requests +sqlalchemy[asyncio] +aiomysql +pydantic