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