Skip to main content

graphwalker_restful/
lib.rs

1pub mod actor;
2mod authoring_dto;
3mod authoring_rest;
4pub mod rest;
5pub mod session;
6pub mod websocket;
7
8use std::net::SocketAddr;
9
10use axum::extract::{FromRef, State, WebSocketUpgrade};
11use axum::response::IntoResponse;
12use axum::routing::{get, post, put};
13use axum::Router;
14use tower_http::cors::CorsLayer;
15
16use crate::actor::spawn_machine_thread;
17use crate::session::SessionManager;
18
19pub use authoring_rest::MAX_AUTHORING_BODY_BYTES;
20
21#[derive(Clone)]
22pub(crate) struct RestApplicationState {
23    execution: rest::RestState,
24    drafts: graphwalker_service::DraftRegistry,
25}
26
27impl FromRef<RestApplicationState> for rest::RestState {
28    fn from_ref(state: &RestApplicationState) -> Self {
29        state.execution.clone()
30    }
31}
32
33impl FromRef<RestApplicationState> for graphwalker_service::DraftRegistry {
34    fn from_ref(state: &RestApplicationState) -> Self {
35        state.drafts.clone()
36    }
37}
38
39pub async fn start_rest_server(
40    port: u16,
41    seed: Option<u64>,
42) -> Result<(), Box<dyn std::error::Error>> {
43    let app = rest_router(seed);
44    serve_rest(port, app).await
45}
46
47pub async fn start_rest_server_with_model(
48    port: u16,
49    seed: Option<u64>,
50    initial_model: Option<String>,
51) -> Result<(), Box<dyn std::error::Error>> {
52    let app = match initial_model {
53        Some(json_body) => {
54            let machine_tx = crate::actor::spawn_machine_thread_with_model(json_body, seed)
55                .map_err(|error| std::io::Error::new(std::io::ErrorKind::InvalidInput, error))?;
56            build_rest_router(RestApplicationState {
57                execution: rest::RestState {
58                    machine_tx,
59                    default_seed: seed,
60                },
61                drafts: graphwalker_service::DraftRegistry::default(),
62            })
63        }
64        None => rest_router(seed),
65    };
66    serve_rest(port, app).await
67}
68
69async fn serve_rest(port: u16, app: Router) -> Result<(), Box<dyn std::error::Error>> {
70    let addr = SocketAddr::from(([0, 0, 0, 0], port));
71    tracing::info!(%addr, "starting REST server");
72    eprintln!("GraphWalker REST server listening on http://{}", addr);
73
74    let listener = tokio::net::TcpListener::bind(addr).await?;
75    axum::serve(listener, app).await?;
76    Ok(())
77}
78
79/// Build the REST application with independent legacy execution and draft state.
80pub fn rest_router(default_seed: Option<u64>) -> Router {
81    rest_router_with_drafts(default_seed, graphwalker_service::DraftRegistry::default())
82}
83
84/// Build the REST application with a caller-supplied draft registry.
85pub fn rest_router_with_drafts(
86    default_seed: Option<u64>,
87    drafts: graphwalker_service::DraftRegistry,
88) -> Router {
89    build_rest_router(RestApplicationState {
90        execution: rest::RestState {
91            machine_tx: spawn_machine_thread(),
92            default_seed,
93        },
94        drafts,
95    })
96}
97
98pub async fn start_websocket_server(port: u16) -> Result<(), Box<dyn std::error::Error>> {
99    let session_mgr = SessionManager::new();
100    let app = build_websocket_router(session_mgr);
101    let addr = SocketAddr::from(([0, 0, 0, 0], port));
102    tracing::info!(%addr, "starting WebSocket server");
103    eprintln!("GraphWalker WebSocket server listening on ws://{}", addr);
104
105    let listener = tokio::net::TcpListener::bind(addr).await?;
106    axum::serve(listener, app).await?;
107    Ok(())
108}
109
110fn build_rest_router(state: RestApplicationState) -> Router {
111    Router::new()
112        .route("/graphwalker/load", post(rest::load))
113        .route("/graphwalker/hasNext", get(rest::has_next))
114        .route("/graphwalker/getNext", get(rest::get_next))
115        .route("/graphwalker/getData", get(rest::get_data))
116        .route("/graphwalker/setData/{script}", put(rest::set_data))
117        .route("/graphwalker/restart", put(rest::restart))
118        .route("/graphwalker/getStatistics", get(rest::get_statistics))
119        .merge(authoring_rest::routes())
120        .layer(CorsLayer::permissive())
121        .with_state(state)
122}
123
124fn build_websocket_router(session_mgr: SessionManager) -> Router {
125    Router::new()
126        .route("/", get(ws_upgrade_handler))
127        .route("/graphwalker", get(ws_upgrade_handler))
128        .layer(CorsLayer::permissive())
129        .with_state(session_mgr)
130}
131
132async fn ws_upgrade_handler(
133    ws: WebSocketUpgrade,
134    State(session_mgr): State<SessionManager>,
135) -> impl IntoResponse {
136    ws.on_upgrade(move |socket| websocket::handle_socket(socket, session_mgr))
137}