use crate::agent::core::{ApiError, Core};
use crate::agent::summary::Summary;
use axum::extract::rejection::JsonRejection;
use axum::extract::{Request, State};
use axum::http::{header, HeaderMap, StatusCode};
use axum::middleware::{self, Next};
use axum::response::{IntoResponse, Response};
use axum::routing::{get, post, put};
use axum::{Json, Router};
use serde::{Deserialize, Deserializer};
use std::sync::Arc;
#[derive(Clone)]
struct Guard {
token: Arc<String>,
port: u16,
}
pub fn router(core: Arc<Core>, token: String, port: u16) -> Router {
let guard = Guard {
token: Arc::new(token),
port,
};
Router::new()
.route("/v1/summary", get(get_summary))
.route("/v1/refresh", post(post_refresh))
.route("/v1/sharing", put(put_sharing))
.route("/v1/workers", put(put_workers))
.route("/v1/price", put(put_price))
.fallback(not_found)
.method_not_allowed_fallback(method_not_allowed)
.layer(middleware::from_fn_with_state(guard, guard_layer))
.with_state(core)
}
fn error(status: u16, code: &str, message: &str) -> Response {
(
StatusCode::from_u16(status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR),
Json(serde_json::json!({ "error": { "code": code, "message": message } })),
)
.into_response()
}
pub fn ct_eq(a: &[u8], b: &[u8]) -> bool {
a.len() == b.len() && a.iter().zip(b).fold(0u8, |acc, (x, y)| acc | (x ^ y)) == 0
}
#[allow(clippy::result_large_err)] pub fn check_request(headers: &HeaderMap, token: &str, port: u16) -> Result<(), Response> {
if headers.contains_key(header::ORIGIN) {
return Err(error(
403,
"forbidden",
"requests from a browser are refused",
));
}
let host = headers
.get(header::HOST)
.and_then(|h| h.to_str().ok())
.unwrap_or("");
if host != format!("127.0.0.1:{port}") && host != format!("localhost:{port}") {
return Err(error(403, "forbidden", "unexpected Host header"));
}
let presented = headers
.get(header::AUTHORIZATION)
.and_then(|h| h.to_str().ok())
.and_then(|h| h.strip_prefix("Bearer "))
.unwrap_or("");
if token.is_empty() || !ct_eq(presented.as_bytes(), token.as_bytes()) {
return Err(error(401, "unauthorized", "missing or invalid agent token"));
}
Ok(())
}
async fn guard_layer(State(g): State<Guard>, req: Request, next: Next) -> Response {
match check_request(req.headers(), &g.token, g.port) {
Ok(()) => next.run(req).await,
Err(resp) => resp,
}
}
fn respond(result: Result<Summary, ApiError>) -> Response {
match result {
Ok(s) => Json(s).into_response(),
Err(e) => error(e.status, e.code, &e.message),
}
}
fn bad_body(e: JsonRejection) -> Response {
error(400, "invalid_input", &e.body_text())
}
async fn get_summary(State(core): State<Arc<Core>>) -> Response {
Json(core.summary()).into_response()
}
async fn post_refresh(State(core): State<Arc<Core>>) -> Response {
core.refresh_hub().await;
core.wake_reconcile.notify_one();
Json(core.summary()).into_response()
}
#[derive(Deserialize)]
struct SharingBody {
on: bool,
}
async fn put_sharing(
State(core): State<Arc<Core>>,
body: Result<Json<SharingBody>, JsonRejection>,
) -> Response {
match body {
Ok(Json(b)) => respond(core.set_sharing(b.on)),
Err(e) => bad_body(e),
}
}
#[derive(Deserialize)]
struct WorkersBody {
count: u32,
}
async fn put_workers(
State(core): State<Arc<Core>>,
body: Result<Json<WorkersBody>, JsonRejection>,
) -> Response {
match body {
Ok(Json(b)) => respond(core.set_workers(b.count)),
Err(e) => bad_body(e),
}
}
fn deserialize_present<'de, D, T>(deserializer: D) -> Result<Option<T>, D::Error>
where
T: Deserialize<'de>,
D: Deserializer<'de>,
{
T::deserialize(deserializer).map(Some)
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct PriceBody {
scope: String,
#[serde(default, deserialize_with = "deserialize_present")]
price_per_hour: Option<Option<f64>>,
}
async fn put_price(
State(core): State<Arc<Core>>,
body: Result<Json<PriceBody>, JsonRejection>,
) -> Response {
match body {
Ok(Json(b)) => match b.price_per_hour {
Some(price) => respond(core.set_price(&b.scope, price).await),
None => error(400, "invalid_input", "price_per_hour is required"),
},
Err(e) => bad_body(e),
}
}
async fn not_found() -> Response {
error(404, "not_found", "no such route")
}
async fn method_not_allowed() -> Response {
error(
405,
"method_not_allowed",
"method not allowed on this route",
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::agent::core::tests::core_with;
use std::time::Duration;
fn headers(pairs: &[(&str, &str)]) -> axum::http::HeaderMap {
let mut h = axum::http::HeaderMap::new();
for (k, v) in pairs {
h.insert(
axum::http::HeaderName::from_bytes(k.as_bytes()).unwrap(),
v.parse().unwrap(),
);
}
h
}
#[test]
fn guards_refuse_browsers_foreign_hosts_and_bad_tokens() {
assert!(check_request(
&headers(&[("host", "127.0.0.1:4720"), ("authorization", "Bearer tok")]),
"tok",
4720
)
.is_ok());
assert!(check_request(
&headers(&[("host", "localhost:4720"), ("authorization", "Bearer tok")]),
"tok",
4720
)
.is_ok());
let status = |pairs: &[(&str, &str)]| {
check_request(&headers(pairs), "tok", 4720)
.unwrap_err()
.status()
.as_u16()
};
assert_eq!(status(&[("host", "127.0.0.1:4720")]), 401, "no token");
assert_eq!(
status(&[("host", "127.0.0.1:4720"), ("authorization", "Bearer nope")]),
401,
"wrong token"
);
assert_eq!(
status(&[("host", "127.0.0.1:4720"), ("authorization", "tok")]),
401,
"Bearer scheme required"
);
assert_eq!(
status(&[
("host", "evil.example:4720"),
("authorization", "Bearer tok")
]),
403,
"DNS rebinding"
);
assert_eq!(
status(&[("host", "127.0.0.1:9999"), ("authorization", "Bearer tok")]),
403
);
assert_eq!(
status(&[
("host", "127.0.0.1:4720"),
("authorization", "Bearer tok"),
("origin", "https://evil.example")
]),
403,
"CSRF"
);
assert!(ct_eq(b"abc", b"abc") && !ct_eq(b"abc", b"abd") && !ct_eq(b"abc", b"ab"));
assert_eq!(
check_request(
&headers(&[("host", "127.0.0.1:4720"), ("authorization", "Bearer ")]),
"",
4720
)
.unwrap_err()
.status()
.as_u16(),
401,
"an empty configured token never authenticates"
);
}
fn serve(core: std::sync::Arc<crate::agent::core::Core>, token: &str) -> u16 {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
listener.set_nonblocking(true).unwrap();
let app = router(core, token.to_string(), port);
std::thread::spawn(move || {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
rt.block_on(async move {
let l = tokio::net::TcpListener::from_std(listener).unwrap();
axum::serve(l, app).await.unwrap();
});
});
port
}
fn call(
method: &str,
port: u16,
path: &str,
token: Option<&str>,
body: &str,
) -> (u16, serde_json::Value, bool) {
let agent = ureq::Agent::new_with_config(
ureq::Agent::config_builder()
.http_status_as_error(false)
.timeout_global(Some(Duration::from_secs(5)))
.build(),
);
let url = format!("http://127.0.0.1:{port}{path}");
let auth = token.map(|t| format!("Bearer {t}")).unwrap_or_default();
let resp = match method {
"GET" => agent
.get(&url)
.header("Authorization", auth.as_str())
.call(),
"PUT" => agent
.put(&url)
.header("Authorization", auth.as_str())
.header("Content-Type", "application/json")
.send(body),
_ => agent
.post(&url)
.header("Authorization", auth.as_str())
.header("Content-Type", "application/json")
.send(body),
}
.unwrap();
let status = resp.status().as_u16();
let cors = resp
.headers()
.keys()
.any(|k| k.as_str().starts_with("access-control"));
let json = resp
.into_body()
.read_json::<serde_json::Value>()
.unwrap_or(serde_json::Value::Null);
(status, json, cors)
}
#[test]
fn the_local_api_round_trip() {
let core = core_with(3);
let port = serve(core, "tok");
let (s, v, cors) = call("GET", port, "/v1/summary", Some("tok"), "");
assert_eq!(
(s, v["v"].as_u64(), v["this_mac"]["workers"]["max"].as_u64()),
(200, Some(1), Some(3))
);
assert!(!cors, "no CORS headers, ever");
let (s, v, _) = call("GET", port, "/v1/summary", None, "");
assert_eq!(
(s, v["error"]["code"].as_str()),
(401, Some("unauthorized"))
);
let (s, v, _) = call("PUT", port, "/v1/workers", Some("tok"), r#"{"count":2}"#);
assert_eq!(
(s, v["this_mac"]["workers"]["desired"].as_u64()),
(200, Some(2))
);
let (s, v, _) = call("PUT", port, "/v1/workers", Some("tok"), r#"{"count":9}"#);
assert_eq!(
(s, v["error"]["code"].as_str()),
(400, Some("invalid_input"))
);
let (s, v, _) = call(
"PUT",
port,
"/v1/workers",
Some("tok"),
r#"{"count":"two"}"#,
);
assert_eq!(
(s, v["error"]["code"].as_str()),
(400, Some("invalid_input"))
);
let (s, v, _) = call("PUT", port, "/v1/sharing", Some("tok"), r#"{"on":true}"#);
assert_eq!((s, v["this_mac"]["sharing"].as_bool()), (200, Some(true)));
let (s, v, _) = call(
"PUT",
port,
"/v1/price",
Some("tok"),
r#"{"scope":"device","price_per_hour":10}"#,
);
assert_eq!(
(s, v["error"]["code"].as_str()),
(503, Some("prerequisite_missing"))
);
let (s, v, _) = call("POST", port, "/v1/refresh", Some("tok"), "");
assert_eq!((s, v["v"].as_u64()), (200, Some(1)));
let (s, _, _) = call("GET", port, "/v1/nope", Some("tok"), "");
assert_eq!(s, 404);
let (s, v, _) = call("GET", port, "/v1/nope", None, "");
assert_eq!(
(s, v["error"]["code"].as_str()),
(401, Some("unauthorized")),
"the guard runs before the 404 too"
);
let (s, v, _) = call(
"PUT",
port,
"/v1/price",
Some("tok"),
r#"{"scope":"device","price":18}"#,
);
assert_eq!(
(s, v["error"]["code"].as_str()),
(400, Some("invalid_input")),
"an unknown field is rejected"
);
let (s, v, _) = call(
"PUT",
port,
"/v1/price",
Some("tok"),
r#"{"scope":"device"}"#,
);
assert_eq!(
(s, v["error"]["code"].as_str()),
(400, Some("invalid_input")),
"price_per_hour is required"
);
let (s, v, _) = call(
"PUT",
port,
"/v1/price",
Some("tok"),
r#"{"scope":"device","price_per_hour":null}"#,
);
assert_ne!(
(s, v["error"]["code"].as_str()),
(400, Some("invalid_input")),
"an explicit null still clears, it's not treated as missing"
);
}
#[test]
fn method_not_allowed_is_json_and_the_guard_runs_first() {
let core = core_with(3);
let port = serve(core, "tok");
let (s, v, _) = call("POST", port, "/v1/summary", Some("tok"), "");
assert_eq!(
(s, v["error"]["code"].as_str()),
(405, Some("method_not_allowed"))
);
let (s, v, _) = call("POST", port, "/v1/summary", None, "");
assert_eq!(
(s, v["error"]["code"].as_str()),
(401, Some("unauthorized")),
"the guard runs before the 405"
);
}
}