synheart_sensor_agent/
server.rs1use axum::{
11 extract::{DefaultBodyLimit, State},
12 http::{header::AUTHORIZATION, HeaderMap, StatusCode},
13 routing::{get, post},
14 Json, Router,
15};
16use serde::{Deserialize, Serialize};
17use std::net::SocketAddr;
18use std::path::PathBuf;
19use std::sync::Arc;
20use tokio::net::TcpListener;
21use tower_http::cors::{AllowOrigin, Any, CorsLayer};
22
23const MAX_BODY_BYTES: usize = 256 * 1024;
26
27#[derive(Debug, Clone)]
29pub struct ServerConfig {
30 pub port: u16,
32 pub state_dir: PathBuf,
34 pub token: Option<String>,
40}
41
42impl ServerConfig {
43 pub fn new(port: u16, state_dir: PathBuf) -> Self {
45 Self {
46 port,
47 state_dir,
48 token: None,
49 }
50 }
51
52 pub fn with_token(mut self, token: Option<String>) -> Self {
54 self.token = token.filter(|t| !t.is_empty());
55 self
56 }
57}
58
59pub struct ServerState {
61 #[allow(dead_code)]
63 state_dir: PathBuf,
64 token: Option<String>,
66}
67
68impl ServerState {
69 pub fn new(config: &ServerConfig) -> Self {
71 Self {
72 state_dir: config.state_dir.clone(),
73 token: config.token.clone(),
74 }
75 }
76}
77
78#[derive(Debug, Clone, Serialize, Deserialize)]
80pub struct BehavioralSession {
81 pub session: serde_json::Value,
83}
84
85#[derive(Debug, Clone, Serialize)]
87pub struct CollectResponse {
88 pub status: String,
90 pub message: String,
92}
93
94#[derive(Serialize)]
96pub struct HealthResponse {
97 pub status: String,
99 pub version: String,
101}
102
103#[derive(Serialize)]
105pub struct ErrorResponse {
106 pub error: String,
108 pub code: String,
110}
111
112async fn health() -> Json<HealthResponse> {
114 Json(HealthResponse {
115 status: "ok".to_string(),
116 version: env!("CARGO_PKG_VERSION").to_string(),
117 })
118}
119
120fn check_auth(
125 state: &ServerState,
126 headers: &HeaderMap,
127) -> Result<(), (StatusCode, Json<ErrorResponse>)> {
128 let Some(expected) = state.token.as_deref() else {
129 return Ok(());
130 };
131
132 let provided = headers
133 .get(AUTHORIZATION)
134 .and_then(|v| v.to_str().ok())
135 .and_then(|v| v.strip_prefix("Bearer "));
136
137 match provided {
138 Some(token) if constant_time_eq(token.as_bytes(), expected.as_bytes()) => Ok(()),
139 _ => Err((
140 StatusCode::UNAUTHORIZED,
141 Json(ErrorResponse {
142 error: "Missing or invalid bearer token".to_string(),
143 code: "UNAUTHORIZED".to_string(),
144 }),
145 )),
146 }
147}
148
149fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
156 if a.len() != b.len() {
157 return false;
158 }
159 let mut diff = 0u8;
160 for (x, y) in a.iter().zip(b.iter()) {
161 diff |= x ^ y;
162 }
163 diff == 0
164}
165
166async fn collect(
172 State(state): State<Arc<ServerState>>,
173 headers: HeaderMap,
174 Json(data): Json<BehavioralSession>,
175) -> Result<Json<CollectResponse>, (StatusCode, Json<ErrorResponse>)> {
176 check_auth(&state, &headers)?;
177
178 if !data.session.is_object() {
179 return Err((
180 StatusCode::BAD_REQUEST,
181 Json(ErrorResponse {
182 error: "Session payload must be a JSON object".to_string(),
183 code: "INVALID_SESSION".to_string(),
184 }),
185 ));
186 }
187
188 Ok(Json(CollectResponse {
189 status: "ok".to_string(),
190 message: "Data accepted".to_string(),
191 }))
192}
193
194pub async fn run(
196 config: ServerConfig,
197) -> anyhow::Result<(SocketAddr, tokio::sync::oneshot::Sender<()>)> {
198 let state = Arc::new(ServerState::new(&config));
199
200 let app = Router::new()
201 .route("/health", get(health))
202 .route("/collect", post(collect))
203 .layer(DefaultBodyLimit::max(MAX_BODY_BYTES))
206 .layer(
207 CorsLayer::new()
214 .allow_origin(AllowOrigin::predicate(|origin, _| {
215 let Ok(origin) = origin.to_str() else {
216 return false;
217 };
218 origin.starts_with("chrome-extension://")
219 || origin == "http://localhost"
220 || origin.starts_with("http://localhost:")
221 || origin == "http://127.0.0.1"
222 || origin.starts_with("http://127.0.0.1:")
223 }))
224 .allow_methods(Any)
225 .allow_headers(Any),
226 )
227 .with_state(state);
228
229 let addr = SocketAddr::from(([127, 0, 0, 1], config.port));
230 let listener = TcpListener::bind(addr).await?;
231 let actual_addr = listener.local_addr()?;
232
233 tracing::info!("Sensor agent server listening on http://{}", actual_addr);
234
235 let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>();
236
237 tokio::spawn(async move {
238 if let Err(e) = axum::serve(listener, app)
239 .with_graceful_shutdown(async {
240 let _ = shutdown_rx.await;
241 tracing::info!("Server shutdown signal received");
242 })
243 .await
244 {
245 tracing::error!("Server error: {}", e);
246 }
247 });
248
249 Ok((actual_addr, shutdown_tx))
250}