use std::fmt;
use std::io::{BufRead, BufReader, Write};
use std::net::{TcpListener, TcpStream, ToSocketAddrs};
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
const MAX_BODY_BYTES: usize = 1 << 20;
const MAX_HEAD_LINE_BYTES: usize = 64 * 1024;
const MAX_HEADER_COUNT: usize = 256;
const READ_TIMEOUT: Duration = Duration::from_secs(30);
use crate::assets::asset_for;
use crate::atelier::AtelierWebState;
use crate::live::{
DEFAULT_PANE, DEFAULT_RESOURCE, DefaultLiveSurfaceFactory, LiveSessionTable,
LiveSurfaceFactory, decode_intent_body, encode_patches, encode_scene, error_json,
};
use sim_kernel::Cx;
use sim_lib_net_core::{CapOutcome, read_capped_line};
use sim_lib_server::{CookbookWebResponse, CookbookWebState};
pub struct ServeConfig {
pub addr: String,
pub atelier_root: PathBuf,
pub dry_run: bool,
pub cookbook: Option<Arc<CookbookWebState>>,
}
impl fmt::Debug for ServeConfig {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ServeConfig")
.field("addr", &self.addr)
.field("atelier_root", &self.atelier_root)
.field("dry_run", &self.dry_run)
.field("cookbook", &self.cookbook.as_ref().map(|_| "<provided>"))
.finish()
}
}
impl Default for ServeConfig {
fn default() -> Self {
Self {
addr: "127.0.0.1:8787".to_owned(),
atelier_root: PathBuf::from(".sim/atelier"),
dry_run: false,
cookbook: None,
}
}
}
pub fn serve_with_cx(cx: &mut Cx, config: &ServeConfig) -> std::io::Result<()> {
serve_with_surface_factory(cx, config, Box::new(DefaultLiveSurfaceFactory))
}
pub fn serve_with_surface_factory(
cx: &mut Cx,
config: &ServeConfig,
surface_factory: Box<dyn LiveSurfaceFactory + Send + Sync>,
) -> std::io::Result<()> {
if config.dry_run {
println!("sim-web-shell: dry-run OK");
return Ok(());
}
let listener = bind(&config.addr)?;
let local = listener.local_addr()?;
let mut state = ShellState::with_surface_factory(config, cx, surface_factory)?;
println!("sim-web-shell: serving shell on http://{local}");
for stream in listener.incoming() {
match stream {
Ok(stream) => {
if let Err(err) = handle(stream, &mut state) {
eprintln!("sim-web-shell: connection error: {err}");
}
}
Err(err) => eprintln!("sim-web-shell: accept error: {err}"),
}
}
Ok(())
}
fn bind(addr: &str) -> std::io::Result<TcpListener> {
let resolved = addr.to_socket_addrs()?.next().ok_or_else(|| {
std::io::Error::new(std::io::ErrorKind::InvalidInput, "no socket address")
})?;
TcpListener::bind(resolved)
}
struct ShellState<'a> {
atelier: AtelierWebState,
cookbook: Arc<CookbookWebState>,
cookbook_cx: &'a mut Cx,
live: LiveSessionTable,
}
impl<'a> ShellState<'a> {
#[cfg(test)]
fn new(config: &ServeConfig, cx: &'a mut Cx) -> std::io::Result<Self> {
Self::with_surface_factory(config, cx, Box::new(DefaultLiveSurfaceFactory))
}
fn with_surface_factory(
config: &ServeConfig,
cx: &'a mut Cx,
surface_factory: Box<dyn LiveSurfaceFactory + Send + Sync>,
) -> std::io::Result<Self> {
Ok(Self {
atelier: AtelierWebState::load(config.atelier_root.clone()),
cookbook: match &config.cookbook {
Some(cookbook) => Arc::clone(cookbook),
None => Arc::new(CookbookWebState::seeded().map_err(io_error)?),
},
cookbook_cx: cx,
live: LiveSessionTable::new(surface_factory),
})
}
}
#[cfg(test)]
pub(crate) fn cookbook_index_for_test(
cx: &mut Cx,
config: &ServeConfig,
) -> std::io::Result<CookbookWebResponse> {
let state = ShellState::new(config, cx)?;
Ok(state
.cookbook
.handle_request("GET", "/api/cookbook", Some(&mut *state.cookbook_cx)))
}
fn io_error(err: impl std::fmt::Display) -> std::io::Error {
std::io::Error::other(err.to_string())
}
fn handle(mut stream: TcpStream, state: &mut ShellState<'_>) -> std::io::Result<()> {
let _ = stream.set_read_timeout(Some(READ_TIMEOUT));
let request = match read_request(&mut stream)? {
ReadOutcome::Request(request) => request,
ReadOutcome::TooLarge => {
write_response(
&mut stream,
413,
"Payload Too Large",
"text/plain; charset=utf-8",
b"payload too large",
)?;
return Ok(());
}
ReadOutcome::Invalid => {
write_response(
&mut stream,
400,
"Bad Request",
"text/plain; charset=utf-8",
b"bad request",
)?;
return Ok(());
}
};
if path_of(&request.target) == "/api/session/intent" {
return write_session_intent(&mut stream, &request, &mut state.live);
}
if path_of(&request.target) == "/api/session/open" {
return write_session_open(&mut stream, &request, &mut state.live);
}
if path_of(&request.target) == "/api/session/close" {
return write_session_close(&mut stream, &request, &mut state.live);
}
if request.target.starts_with("/api/cookbook") {
let response = state.cookbook.handle_request(
&request.method,
&request.target,
Some(&mut *state.cookbook_cx),
);
return write_cookbook_response(&mut stream, &response);
}
if let Some(response) = state.atelier.response(&request.method, &request.target) {
return write_response(
&mut stream,
response.status,
status_text(response.status),
response.content_type,
response.body.as_bytes(),
);
}
if request.method != "GET" {
write_response(
&mut stream,
405,
"Method Not Allowed",
"text/plain; charset=utf-8",
b"method not allowed",
)?;
return Ok(());
}
match asset_for(&request.target) {
Some(asset) => write_response(&mut stream, 200, "OK", asset.content_type, asset.body),
None => write_response(
&mut stream,
404,
"Not Found",
"text/plain; charset=utf-8",
b"not found",
),
}
}
#[derive(Debug)]
struct RequestLine {
method: String,
target: String,
body: String,
}
#[derive(Debug)]
enum ReadOutcome {
Request(RequestLine),
TooLarge,
Invalid,
}
fn read_request(stream: &mut TcpStream) -> std::io::Result<ReadOutcome> {
let mut reader = BufReader::new(stream);
read_request_from(&mut reader)
}
fn read_request_from(reader: &mut impl BufRead) -> std::io::Result<ReadOutcome> {
let mut request_line = String::new();
match read_capped_line(reader, &mut request_line, MAX_HEAD_LINE_BYTES)? {
CapOutcome::TooLarge => return Ok(ReadOutcome::TooLarge),
CapOutcome::Eof => return Ok(ReadOutcome::Invalid),
CapOutcome::Line => {}
}
let mut content_length = 0usize;
let mut header = String::new();
let mut header_count = 0usize;
loop {
header_count += 1;
if header_count > MAX_HEADER_COUNT {
return Ok(ReadOutcome::TooLarge);
}
match read_capped_line(reader, &mut header, MAX_HEAD_LINE_BYTES)? {
CapOutcome::TooLarge => return Ok(ReadOutcome::TooLarge),
CapOutcome::Eof => break,
CapOutcome::Line => {}
}
if header == "\r\n" || header == "\n" {
break;
}
if let Some((name, value)) = header.split_once(':')
&& name.trim().eq_ignore_ascii_case("content-length")
{
content_length = value.trim().parse().unwrap_or(0);
}
}
if content_length > MAX_BODY_BYTES {
return Ok(ReadOutcome::TooLarge);
}
let mut body = vec![0u8; content_length];
if content_length > 0 {
reader.read_exact(&mut body)?;
}
let body = String::from_utf8_lossy(&body).into_owned();
let mut parts = request_line.split_whitespace();
let method = parts.next();
let target = parts.next();
match (method, target) {
(Some(method @ ("GET" | "POST")), Some(target)) => Ok(ReadOutcome::Request(RequestLine {
method: method.to_owned(),
target: target.to_owned(),
body,
})),
_ => Ok(ReadOutcome::Invalid),
}
}
fn write_session_intent(
stream: &mut (impl Write + ?Sized),
request: &RequestLine,
live: &mut LiveSessionTable,
) -> std::io::Result<()> {
if request.method != "POST" {
return write_json(stream, 405, &error_json("intent route requires POST"));
}
let session_id = match query_value(&request.target, "session") {
Ok(Some(value)) => value,
Ok(None) => return write_json(stream, 400, &error_json("missing session id")),
Err(err) => return write_json(stream, 400, &error_json(&err.to_string())),
};
let pane = match query_value(&request.target, "pane") {
Ok(Some(value)) => value,
Ok(None) => DEFAULT_PANE.to_owned(),
Err(err) => return write_json(stream, 400, &error_json(&err.to_string())),
};
let intent = match decode_intent_body(&request.body) {
Ok(intent) => intent,
Err(err) => return write_json(stream, 400, &error_json(&err)),
};
match live.submit(&session_id, &pane, &intent) {
Ok(updates) => write_json(stream, 200, &encode_patches(&updates)),
Err(err) => write_json(stream, 400, &error_json(&err.to_string())),
}
}
fn write_session_open(
stream: &mut (impl Write + ?Sized),
request: &RequestLine,
live: &mut LiveSessionTable,
) -> std::io::Result<()> {
if request.method != "GET" {
return write_json(stream, 405, &error_json("open route requires GET"));
}
let session_id = match query_value(&request.target, "session") {
Ok(value) => value,
Err(err) => return write_json(stream, 400, &error_json(&err.to_string())),
};
let resource = match query_value(&request.target, "resource") {
Ok(Some(value)) => value,
Ok(None) => DEFAULT_RESOURCE.to_owned(),
Err(err) => return write_json(stream, 400, &error_json(&err.to_string())),
};
let pane = match query_value(&request.target, "pane") {
Ok(Some(value)) => value,
Ok(None) => DEFAULT_PANE.to_owned(),
Err(err) => return write_json(stream, 400, &error_json(&err.to_string())),
};
match live.open(session_id.as_deref(), &resource, &pane) {
Ok((session_id, scene)) => {
write_json(stream, 200, &encode_session_open(&session_id, &scene))
}
Err(err) => write_json(stream, 400, &error_json(&err.to_string())),
}
}
fn write_session_close(
stream: &mut (impl Write + ?Sized),
request: &RequestLine,
live: &mut LiveSessionTable,
) -> std::io::Result<()> {
if request.method != "POST" {
return write_json(stream, 405, &error_json("close route requires POST"));
}
let session_id = match query_value(&request.target, "session") {
Ok(Some(value)) => value,
Ok(None) => return write_json(stream, 400, &error_json("missing session id")),
Err(err) => return write_json(stream, 400, &error_json(&err.to_string())),
};
match live.close(&session_id) {
Ok(()) => write_json(stream, 200, r#"{"ok":true}"#),
Err(err) => write_json(stream, 400, &error_json(&err)),
}
}
fn encode_session_open(session_id: &str, scene: &sim_kernel::Expr) -> String {
let mut value: serde_json::Value =
serde_json::from_str(&encode_scene(scene)).expect("encode_scene emits JSON object");
if let Some(object) = value.as_object_mut() {
object.insert(
"session".to_owned(),
serde_json::Value::String(session_id.to_owned()),
);
}
value.to_string()
}
fn path_of(target: &str) -> &str {
target.split(['?', '#']).next().unwrap_or(target)
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct QueryError {
key: String,
reason: String,
}
impl QueryError {
fn new(key: &str, reason: impl Into<String>) -> Self {
Self {
key: key.to_owned(),
reason: reason.into(),
}
}
}
impl fmt::Display for QueryError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"malformed query value for '{}': {}",
self.key, self.reason
)
}
}
fn query_value(target: &str, key: &str) -> Result<Option<String>, QueryError> {
let Some((_, query_and_fragment)) = target.split_once('?') else {
return Ok(None);
};
let query = query_and_fragment
.split('#')
.next()
.unwrap_or(query_and_fragment);
for pair in query.split('&') {
let (name, value) = pair.split_once('=').unwrap_or((pair, ""));
if name == key {
return percent_decode_query(value)
.map(Some)
.map_err(|reason| QueryError::new(key, reason));
}
}
Ok(None)
}
fn percent_decode_query(value: &str) -> Result<String, String> {
let bytes = value.as_bytes();
let mut decoded = Vec::with_capacity(bytes.len());
let mut index = 0usize;
while index < bytes.len() {
match bytes[index] {
b'%' => {
if index + 2 >= bytes.len() {
return Err("incomplete percent escape".to_owned());
}
let high = hex_digit(bytes[index + 1])
.ok_or_else(|| "invalid percent escape".to_owned())?;
let low = hex_digit(bytes[index + 2])
.ok_or_else(|| "invalid percent escape".to_owned())?;
decoded.push((high << 4) | low);
index += 3;
}
b'+' => {
decoded.push(b' ');
index += 1;
}
byte => {
decoded.push(byte);
index += 1;
}
}
}
String::from_utf8(decoded).map_err(|_| "decoded value is not UTF-8".to_owned())
}
fn hex_digit(byte: u8) -> Option<u8> {
match byte {
b'0'..=b'9' => Some(byte - b'0'),
b'a'..=b'f' => Some(byte - b'a' + 10),
b'A'..=b'F' => Some(byte - b'A' + 10),
_ => None,
}
}
fn write_json(stream: &mut (impl Write + ?Sized), status: u16, body: &str) -> std::io::Result<()> {
write_response(
stream,
status,
status_text(status),
"application/json; charset=utf-8",
body.as_bytes(),
)
}
fn write_cookbook_response(
stream: &mut (impl Write + ?Sized),
response: &CookbookWebResponse,
) -> std::io::Result<()> {
write_response(
stream,
response.status,
status_text(response.status),
response.content_type,
response.body.as_bytes(),
)
}
fn write_response(
stream: &mut (impl Write + ?Sized),
status: u16,
reason: &str,
content_type: &str,
body: &[u8],
) -> std::io::Result<()> {
let header = format!(
"HTTP/1.1 {status} {reason}\r\nContent-Type: {content_type}\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
stream.write_all(header.as_bytes())?;
stream.write_all(body)?;
stream.flush()
}
fn status_text(status: u16) -> &'static str {
match status {
200 => "OK",
201 => "Created",
204 => "No Content",
301 => "Moved Permanently",
302 => "Found",
304 => "Not Modified",
400 => "Bad Request",
401 => "Unauthorized",
403 => "Forbidden",
404 => "Not Found",
405 => "Method Not Allowed",
409 => "Conflict",
413 => "Payload Too Large",
422 => "Unprocessable Entity",
429 => "Too Many Requests",
500 => "Internal Server Error",
501 => "Not Implemented",
503 => "Service Unavailable",
other => match other / 100 {
1 => "Informational",
2 => "OK",
3 => "Redirection",
4 => "Client Error",
_ => "Internal Server Error",
},
}
}
#[cfg(test)]
#[path = "serve_tests.rs"]
mod tests;