eidetica/sync/transports/
http.rs1use std::{net::SocketAddr, sync::Arc, time::Duration};
7
8use async_trait::async_trait;
9use axum::{
10 Router,
11 extract::{ConnectInfo, Json as ExtractJson, State},
12 response::Json,
13 routing::post,
14};
15use serde::{Deserialize, Serialize};
16use tokio::sync::oneshot;
17
18use super::{SyncTransport, TransportBuilder, TransportConfig, shared::*};
19use crate::{
20 Result,
21 crdt::Doc,
22 store::Registered,
23 sync::{
24 error::{SyncError, TimeoutPhase},
25 handler::SyncHandler,
26 peer_types::Address,
27 protocol::{RequestContext, SyncRequest, SyncResponse},
28 },
29};
30
31const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
36
37const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
51
52fn build_client() -> Result<reqwest::Client> {
54 reqwest::Client::builder()
55 .connect_timeout(CONNECT_TIMEOUT)
56 .read_timeout(REQUEST_TIMEOUT)
57 .build()
58 .map_err(|e| SyncError::Network(format!("Failed to build HTTP client: {e}")).into())
59}
60
61fn timeout_error(error: &reqwest::Error, address: &str) -> Option<SyncError> {
69 if !error.is_timeout() {
70 return None;
71 }
72
73 let (phase, elapsed) = if error.is_connect() {
74 (TimeoutPhase::Connect, CONNECT_TIMEOUT)
75 } else {
76 (TimeoutPhase::Request, REQUEST_TIMEOUT)
77 };
78
79 Some(SyncError::Timeout {
80 address: address.to_string(),
81 phase,
82 elapsed,
83 })
84}
85
86#[derive(Debug, Clone, Serialize, Deserialize, Default)]
90pub struct HttpTransportConfig {
91 #[serde(default, skip_serializing_if = "Option::is_none")]
94 pub bind_address: Option<String>,
95}
96
97impl Registered for HttpTransportConfig {
98 fn type_id() -> &'static str {
99 "http:v0"
100 }
101}
102
103impl TransportConfig for HttpTransportConfig {}
104
105#[derive(Debug, Clone, Default)]
120pub struct HttpTransportBuilder {
121 bind_address: Option<String>,
122}
123
124impl HttpTransportBuilder {
125 pub fn new() -> Self {
127 Self::default()
128 }
129
130 pub fn bind(mut self, addr: impl Into<String>) -> Self {
142 self.bind_address = Some(addr.into());
143 self
144 }
145
146 pub fn build_sync(self) -> Result<HttpTransport> {
151 Ok(HttpTransport {
152 server_state: ServerState::new(),
153 bind_address: self.bind_address,
154 client: build_client()?,
155 })
156 }
157}
158
159#[async_trait]
160impl TransportBuilder for HttpTransportBuilder {
161 type Transport = HttpTransport;
162
163 async fn build(self, _persisted: Doc) -> Result<(Self::Transport, Option<Doc>)> {
167 let transport = HttpTransport {
168 server_state: ServerState::new(),
169 bind_address: self.bind_address,
170 client: build_client()?,
171 };
172 Ok((transport, None))
173 }
174}
175
176pub struct HttpTransport {
178 server_state: ServerState,
180 bind_address: Option<String>,
182 client: reqwest::Client,
185}
186
187impl HttpTransport {
188 pub const TRANSPORT_TYPE: &'static str = "http";
190
191 pub fn new() -> Result<Self> {
193 Ok(Self {
194 server_state: ServerState::new(),
195 bind_address: None,
196 client: build_client()?,
197 })
198 }
199
200 pub fn builder() -> HttpTransportBuilder {
202 HttpTransportBuilder::new()
203 }
204
205 fn create_router(handler: Arc<dyn SyncHandler>) -> Router {
207 Router::new()
208 .route("/api/v0", post(handle_sync_request))
209 .with_state(handler)
210 }
211}
212
213#[async_trait]
214impl SyncTransport for HttpTransport {
215 fn transport_type(&self) -> &'static str {
216 Self::TRANSPORT_TYPE
217 }
218
219 fn can_handle_address(&self, address: &Address) -> bool {
220 address.transport_type == Self::TRANSPORT_TYPE
221 }
222
223 async fn start_server(&self, handler: Arc<dyn SyncHandler>) -> Result<()> {
224 let Some(effective_addr) = self.bind_address.as_deref() else {
226 return Ok(());
227 };
228
229 let start = self.server_state.begin_start(effective_addr)?;
230
231 let socket_addr: SocketAddr =
232 effective_addr.parse().map_err(|e| SyncError::ServerBind {
233 address: effective_addr.to_string(),
234 reason: format!("Invalid address: {e}"),
235 })?;
236
237 let router = Self::create_router(handler);
238
239 let (ready_tx, ready_rx) = oneshot::channel();
241 let (shutdown_tx, shutdown_rx) = oneshot::channel();
242
243 let (addr_tx, addr_rx) = oneshot::channel::<SocketAddr>();
245
246 let listener = match tokio::net::TcpListener::bind(socket_addr).await {
247 Ok(listener) => listener,
248 Err(error) => {
249 return Err(SyncError::ServerBind {
250 address: effective_addr.to_string(),
251 reason: error.to_string(),
252 }
253 .into());
254 }
255 };
256 let actual_addr = match listener.local_addr() {
257 Ok(address) => address,
258 Err(error) => {
259 return Err(SyncError::ServerBind {
260 address: effective_addr.to_string(),
261 reason: error.to_string(),
262 }
263 .into());
264 }
265 };
266
267 tokio::spawn(async move {
269 let _ = addr_tx.send(actual_addr);
271
272 let _ = ready_tx.send(());
274
275 if let Err(error) = axum::serve(
278 listener,
279 router.into_make_service_with_connect_info::<SocketAddr>(),
280 )
281 .with_graceful_shutdown(async move {
282 let _ = shutdown_rx.await;
283 })
284 .await
285 {
286 tracing::error!(%error, "HTTP sync server failed");
287 }
288 });
289
290 let actual_addr = match addr_rx.await {
292 Ok(address) => address,
293 Err(_) => {
294 return Err(SyncError::ServerBind {
295 address: effective_addr.to_string(),
296 reason: "Failed to get actual server address".to_string(),
297 }
298 .into());
299 }
300 };
301
302 wait_for_ready(ready_rx, effective_addr).await?;
304
305 start.complete(actual_addr.to_string(), shutdown_tx);
307
308 Ok(())
309 }
310
311 async fn stop_server(&self) -> Result<()> {
312 if !self.server_state.is_running() {
313 return Err(SyncError::ServerNotRunning.into());
314 }
315
316 self.server_state.stop_server();
318
319 Ok(())
320 }
321
322 async fn send_request(&self, address: &Address, request: &SyncRequest) -> Result<SyncResponse> {
323 if !self.can_handle_address(address) {
324 return Err(SyncError::UnsupportedTransport {
325 transport_type: address.transport_type.clone(),
326 }
327 .into());
328 }
329
330 let url = format!("http://{}/api/v0", address.address);
331
332 let response = self
333 .client
334 .post(&url)
335 .json(&request) .send()
337 .await
338 .map_err(|e| {
339 timeout_error(&e, &address.address).unwrap_or_else(|| SyncError::ConnectionFailed {
340 address: address.address.clone(),
341 reason: e.to_string(),
342 })
343 })?;
344
345 if !response.status().is_success() {
346 return Err(SyncError::Network(format!(
347 "Server returned error: {}",
348 response.status()
349 ))
350 .into());
351 }
352
353 let sync_response: SyncResponse = response.json().await.map_err(|e| {
358 timeout_error(&e, &address.address)
359 .unwrap_or_else(|| SyncError::Network(format!("Failed to parse response: {e}")))
360 })?;
361
362 Ok(sync_response)
363 }
364
365 fn is_server_running(&self) -> bool {
366 self.server_state.is_running()
367 }
368
369 fn get_server_address(&self) -> Result<String> {
370 self.server_state.get_address().map_err(|e| e.into())
371 }
372}
373
374async fn handle_sync_request(
376 State(handler): State<Arc<dyn SyncHandler>>,
377 ConnectInfo(addr): ConnectInfo<SocketAddr>,
378 ExtractJson(request): ExtractJson<SyncRequest>,
379) -> Json<SyncResponse> {
380 let peer_pubkey = match &request {
382 SyncRequest::SyncTree(sync_tree_request) => sync_tree_request.peer_pubkey.clone(),
383 _ => None,
384 };
385
386 let context = RequestContext {
388 remote_address: Some(Address {
389 transport_type: HttpTransport::TRANSPORT_TYPE.to_string(),
390 address: addr.to_string(),
391 }),
392 peer_pubkey,
393 };
394
395 let response = handler.handle_request(&request, &context).await;
397
398 Json(response)
399}
400
401#[cfg(test)]
402mod tests {
403 use std::future::pending;
404
405 use tokio::net::TcpListener;
406
407 use super::*;
408
409 #[tokio::test]
415 async fn a_peer_that_accepts_and_says_nothing_times_out_in_the_request_phase() {
416 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
417 let address = listener.local_addr().unwrap().to_string();
418
419 tokio::spawn(async move {
421 let _accepted = listener.accept().await;
422 pending::<()>().await;
423 });
424
425 let client = reqwest::Client::builder()
426 .read_timeout(Duration::from_millis(200))
427 .build()
428 .unwrap();
429 let error = client
430 .post(format!("http://{address}/api/v0"))
431 .send()
432 .await
433 .expect_err("a peer that never answers cannot produce a response");
434
435 match timeout_error(&error, &address) {
436 Some(SyncError::Timeout { phase, .. }) => assert_eq!(phase, TimeoutPhase::Request),
437 other => panic!("expected a request-phase timeout, got {other:?}"),
438 }
439 }
440
441 #[tokio::test]
445 async fn a_refused_connection_is_not_reported_as_a_timeout() {
446 let address = {
448 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
449 listener.local_addr().unwrap().to_string()
450 };
451
452 let client = build_client().unwrap();
453 let error = client
454 .post(format!("http://{address}/api/v0"))
455 .send()
456 .await
457 .expect_err("nothing is listening on that port");
458
459 assert!(timeout_error(&error, &address).is_none());
460 }
461}