pub mod actor;
mod authoring_dto;
mod authoring_rest;
pub mod rest;
pub mod session;
pub mod websocket;
use std::net::SocketAddr;
use axum::extract::{FromRef, State, WebSocketUpgrade};
use axum::response::IntoResponse;
use axum::routing::{get, post, put};
use axum::Router;
use tower_http::cors::CorsLayer;
use crate::actor::spawn_machine_thread;
use crate::session::SessionManager;
pub use authoring_rest::MAX_AUTHORING_BODY_BYTES;
#[derive(Clone)]
pub(crate) struct RestApplicationState {
execution: rest::RestState,
drafts: graphwalker_service::DraftRegistry,
}
impl FromRef<RestApplicationState> for rest::RestState {
fn from_ref(state: &RestApplicationState) -> Self {
state.execution.clone()
}
}
impl FromRef<RestApplicationState> for graphwalker_service::DraftRegistry {
fn from_ref(state: &RestApplicationState) -> Self {
state.drafts.clone()
}
}
pub async fn start_rest_server(
port: u16,
seed: Option<u64>,
) -> Result<(), Box<dyn std::error::Error>> {
let app = rest_router(seed);
serve_rest(port, app).await
}
pub async fn start_rest_server_with_model(
port: u16,
seed: Option<u64>,
initial_model: Option<String>,
) -> Result<(), Box<dyn std::error::Error>> {
let app = match initial_model {
Some(json_body) => {
let machine_tx = crate::actor::spawn_machine_thread_with_model(json_body, seed)
.map_err(|error| std::io::Error::new(std::io::ErrorKind::InvalidInput, error))?;
build_rest_router(RestApplicationState {
execution: rest::RestState {
machine_tx,
default_seed: seed,
},
drafts: graphwalker_service::DraftRegistry::default(),
})
}
None => rest_router(seed),
};
serve_rest(port, app).await
}
async fn serve_rest(port: u16, app: Router) -> Result<(), Box<dyn std::error::Error>> {
let addr = SocketAddr::from(([0, 0, 0, 0], port));
eprintln!("GraphWalker REST server listening on http://{}", addr);
let listener = tokio::net::TcpListener::bind(addr).await?;
axum::serve(listener, app).await?;
Ok(())
}
pub fn rest_router(default_seed: Option<u64>) -> Router {
rest_router_with_drafts(default_seed, graphwalker_service::DraftRegistry::default())
}
pub fn rest_router_with_drafts(
default_seed: Option<u64>,
drafts: graphwalker_service::DraftRegistry,
) -> Router {
build_rest_router(RestApplicationState {
execution: rest::RestState {
machine_tx: spawn_machine_thread(),
default_seed,
},
drafts,
})
}
pub async fn start_websocket_server(port: u16) -> Result<(), Box<dyn std::error::Error>> {
let session_mgr = SessionManager::new();
let app = build_websocket_router(session_mgr);
let addr = SocketAddr::from(([0, 0, 0, 0], port));
eprintln!("GraphWalker WebSocket server listening on ws://{}", addr);
let listener = tokio::net::TcpListener::bind(addr).await?;
axum::serve(listener, app).await?;
Ok(())
}
fn build_rest_router(state: RestApplicationState) -> Router {
Router::new()
.route("/graphwalker/load", post(rest::load))
.route("/graphwalker/hasNext", get(rest::has_next))
.route("/graphwalker/getNext", get(rest::get_next))
.route("/graphwalker/getData", get(rest::get_data))
.route("/graphwalker/setData/{script}", put(rest::set_data))
.route("/graphwalker/restart", put(rest::restart))
.route("/graphwalker/getStatistics", get(rest::get_statistics))
.merge(authoring_rest::routes())
.layer(CorsLayer::permissive())
.with_state(state)
}
fn build_websocket_router(session_mgr: SessionManager) -> Router {
Router::new()
.route("/", get(ws_upgrade_handler))
.route("/graphwalker", get(ws_upgrade_handler))
.layer(CorsLayer::permissive())
.with_state(session_mgr)
}
async fn ws_upgrade_handler(
ws: WebSocketUpgrade,
State(session_mgr): State<SessionManager>,
) -> impl IntoResponse {
ws.on_upgrade(move |socket| websocket::handle_socket(socket, session_mgr))
}