Skip to main content

hibp_sync_client/
sync.rs

1use std::io;
2use std::path::{Path, PathBuf};
3use std::time::Duration;
4
5use chrono::{DateTime, SecondsFormat, Utc};
6use futures_util::StreamExt;
7use serde::{Deserialize, Serialize};
8use tokio::fs;
9
10use crate::client::Client;
11use crate::error::Error;
12use crate::wire::decode_segment_stream;
13
14const FIRST_BYTE_TIMEOUT: Duration = Duration::from_secs(30);
15const IDLE_STREAM_TIMEOUT: Duration = Duration::from_secs(120);
16const BUSY_RETRY_DELAY: Duration = Duration::from_secs(30);
17const MAX_BUSY_RETRIES: u32 = 480;
18
19pub struct Config {
20    pub server_url: http::Uri,
21    pub data_dir: PathBuf,
22    pub segments: u8,
23}
24
25pub enum Outcome {
26    UpToDate,
27    DeltaSync { changed_count: usize },
28    FullSync { file_count: usize },
29}
30
31#[derive(Serialize, Deserialize, Default)]
32struct LocalState {
33    last_updated: Option<DateTime<Utc>>,
34}
35
36#[derive(Debug, Serialize, Deserialize)]
37struct Plan {
38    server_last_updated: DateTime<Utc>,
39    since: Option<String>,
40    segments: u8,
41}
42
43#[tracing::instrument(skip_all)]
44pub async fn sync(config: &Config) -> Result<Outcome, Error> {
45    tracing::info!(server_url = %config.server_url, segments = config.segments, "starting sync");
46
47    if config.segments == 0 {
48        return Err(Error::InvalidConfig("segments must be >= 1"));
49    }
50
51    let staging = config.data_dir.join(".staging");
52    let complete_marker = staging.join(".complete");
53    let plan_path = staging.join(".sync-plan.json");
54
55    if staging.exists() {
56        if complete_marker.exists() {
57            tracing::info!("staging/.complete exists - finishing interrupted commit");
58            return finish_commit(&staging, &config.data_dir).await;
59        } else if plan_path.exists() {
60            tracing::info!("resuming interrupted download");
61            let plan: Plan = serde_json::from_slice(&fs::read(&plan_path).await?)?;
62
63            let client = Client::new(&config.server_url)?;
64            let status = client.status().await?;
65
66            if status.last_updated != Some(plan.server_last_updated) {
67                tracing::warn!(
68                    "server state changed since last attempt; discarding staging and starting fresh"
69                );
70                clear_staging(&staging).await?;
71            } else {
72                fetch_missing_segments(&config.server_url, &staging, &plan).await?;
73                fs::write(&complete_marker, b"").await?;
74                return finish_commit(&staging, &config.data_dir).await;
75            }
76        } else {
77            tracing::warn!("staging exists without .sync-plan.json; discarding");
78            clear_staging(&staging).await?;
79        }
80    }
81
82    let state_path = config.data_dir.join("sync-state.json");
83    let local: LocalState = match fs::read(&state_path).await {
84        Ok(bytes) => serde_json::from_slice(&bytes)?,
85        Err(e) if e.kind() == io::ErrorKind::NotFound => LocalState::default(),
86        Err(e) => return Err(e.into()),
87    };
88    let client = Client::new(&config.server_url)?;
89    let server_last_updated = match client.status().await?.last_updated {
90        Some(t) => t,
91        None => return Ok(Outcome::UpToDate),
92    };
93
94    if Some(server_last_updated) <= local.last_updated {
95        return Ok(Outcome::UpToDate);
96    }
97
98    // Use Z-suffix format (e.g. "2026-01-01T00:00:00Z") so the value is URL-safe
99    // without encoding when used as a query parameter.
100    let since_opt: Option<String> = if local.last_updated.is_none() {
101        None
102    } else {
103        let changed = client.changed().await?;
104        if changed.prev_last_updated == local.last_updated {
105            local.last_updated.map(|t| t.to_rfc3339_opts(SecondsFormat::Secs, true))
106        } else {
107            tracing::warn!(
108                "server prev_last_updated does not match local last_updated; falling back to full sync"
109            );
110            None
111        }
112    };
113
114    fs::create_dir_all(&staging).await?;
115
116    let plan = Plan { server_last_updated, since: since_opt, segments: config.segments };
117    fs::write(&plan_path, serde_json::to_vec_pretty(&plan)?).await?;
118
119    fetch_missing_segments(&config.server_url, &staging, &plan).await?;
120    fs::write(&complete_marker, b"").await?;
121
122    finish_commit(&staging, &config.data_dir).await
123}
124
125#[tracing::instrument(skip(server_url, staging), fields(segments = plan.segments, since = plan.since.as_deref()))]
126async fn fetch_missing_segments(
127    server_url: &http::Uri,
128    staging: &Path,
129    plan: &Plan,
130) -> Result<(), Error> {
131    let segments = plan.segments;
132    let client = Client::new(server_url)?;
133
134    let already_done = (0..segments)
135        .filter(|&s| staging.join(format!(".seg.{}.done", s)).exists())
136        .count();
137    tracing::info!(
138        total = segments,
139        already_done,
140        remaining = segments as usize - already_done,
141        "fetching segments"
142    );
143
144    for seg in 0..segments {
145        if staging.join(format!(".seg.{}.done", seg)).exists() {
146            tracing::debug!(
147                segment = seg + 1,
148                of = segments,
149                "already downloaded, skipping"
150            );
151            continue;
152        }
153        tracing::info!(segment = seg + 1, of = segments, "downloading segment");
154        let mut busy_retries = 0u32;
155        loop {
156            match fetch_segment_with_retry(&client, seg, segments, plan.since.as_deref(), staging)
157                .await
158            {
159                Ok(()) => break,
160                Err(Error::ServerBusy) if busy_retries < MAX_BUSY_RETRIES => {
161                    busy_retries += 1;
162                    let jitter_range_ms = (BUSY_RETRY_DELAY.as_millis() as u64 / 4).max(1);
163                    let noise_ms = std::time::SystemTime::now()
164                        .duration_since(std::time::UNIX_EPOCH)
165                        .unwrap_or_default()
166                        .subsec_millis() as u64
167                        % (jitter_range_ms * 2);
168                    let delay = BUSY_RETRY_DELAY
169                        .saturating_sub(Duration::from_millis(jitter_range_ms))
170                        + Duration::from_millis(noise_ms);
171                    tracing::warn!(
172                        attempt = busy_retries,
173                        max = MAX_BUSY_RETRIES,
174                        delay_secs = delay.as_secs(),
175                        "server busy, waiting to retry"
176                    );
177                    tokio::time::sleep(delay).await;
178                }
179                Err(e) => return Err(e),
180            }
181        }
182    }
183
184    Ok(())
185}
186
187#[tracing::instrument(skip(client, staging))]
188async fn fetch_segment_with_retry(
189    client: &Client,
190    segment: u8,
191    of: u8,
192    since: Option<&str>,
193    staging: &Path,
194) -> Result<(), Error> {
195    const MAX_RETRIES: u32 = 5;
196    let mut delay = Duration::from_millis(500);
197    let mut last_result = Ok(());
198
199    for attempt in 0..MAX_RETRIES {
200        if attempt > 0 {
201            let jitter_range_ms = (delay.as_millis() as u64 / 4).max(1);
202            let noise_ms = std::time::SystemTime::now()
203                .duration_since(std::time::UNIX_EPOCH)
204                .unwrap_or_default()
205                .subsec_millis() as u64
206                % (jitter_range_ms * 2);
207            let jittered = delay.saturating_sub(Duration::from_millis(jitter_range_ms))
208                + Duration::from_millis(noise_ms);
209            tracing::warn!(
210                attempt,
211                error = %last_result.as_ref().unwrap_err(),
212                delay_ms = jittered.as_millis(),
213                "segment fetch failed, retrying"
214            );
215            tokio::time::sleep(jittered).await;
216            delay *= 2;
217        }
218        let result = match client.segment_stream(segment, of, since).await {
219            Ok(decoder) => try_stream_segment(decoder, staging).await,
220            Err(Error::ServerBusy) => return Err(Error::ServerBusy),
221            Err(e) => Err(e),
222        };
223        match result {
224            Ok(files_written) => {
225                tracing::info!(files = files_written, "segment complete");
226                fs::write(staging.join(format!(".seg.{}.done", segment)), b"").await?;
227                return Ok(());
228            }
229            Err(e) => last_result = Err(e),
230        }
231    }
232
233    tracing::error!(error = %last_result.as_ref().unwrap_err(), "all retries exhausted");
234    last_result
235}
236
237async fn try_stream_segment<R: tokio::io::AsyncRead + Unpin + Send + 'static>(
238    decoder: R,
239    staging: &Path,
240) -> Result<usize, Error> {
241    let mut stream = Box::pin(decode_segment_stream(decoder));
242    let mut files_written: usize = 0;
243    let mut first = true;
244
245    loop {
246        let (timeout, timeout_err) = if first {
247            (FIRST_BYTE_TIMEOUT, "first byte")
248        } else {
249            (IDLE_STREAM_TIMEOUT, "idle stream read")
250        };
251        let Some(entry_res) = tokio::time::timeout(timeout, stream.next())
252            .await
253            .map_err(|_| Error::Timeout(timeout_err))?
254        else {
255            break;
256        };
257        first = false;
258        let entry = entry_res?;
259        let prefix_str = std::str::from_utf8(&entry.prefix)
260            .map_err(|e| Error::Decode(format!("invalid prefix bytes: {e}")))?;
261        fs::write(staging.join(format!("{}.bin", prefix_str)), &entry.content).await?;
262        files_written += 1;
263    }
264
265    Ok(files_written)
266}
267
268#[tracing::instrument(skip_all)]
269async fn finish_commit(staging: &Path, data_dir: &Path) -> Result<Outcome, Error> {
270    let plan: Plan = serde_json::from_slice(&fs::read(staging.join(".sync-plan.json")).await?)?;
271
272    let mut entries = fs::read_dir(staging).await?;
273    let mut file_count = 0usize;
274    while let Some(entry) = entries.next_entry().await? {
275        let src = entry.path();
276        if src.extension().is_some_and(|e| e == "bin") {
277            fs::rename(&src, data_dir.join(src.file_name().unwrap())).await?;
278            file_count += 1;
279        }
280    }
281
282    let state_path = data_dir.join("sync-state.json");
283    let new_state = LocalState { last_updated: Some(plan.server_last_updated) };
284    let tmp = state_path.with_extension("json.tmp");
285    fs::write(&tmp, serde_json::to_vec_pretty(&new_state)?).await?;
286    fs::rename(&tmp, &state_path).await?;
287
288    clear_staging(staging).await?;
289
290    Ok(match plan.since {
291        Some(_) => Outcome::DeltaSync { changed_count: file_count },
292        None => Outcome::FullSync { file_count },
293    })
294}
295
296async fn clear_staging(staging: &Path) -> Result<(), Error> {
297    fs::remove_dir_all(staging).await?;
298    Ok(())
299}
300
301#[cfg(test)]
302mod tests {
303    use super::*;
304
305    #[tokio::test]
306    async fn sync_rejects_zero_segments() {
307        let tmp = tempfile::tempdir().unwrap();
308        let cfg = Config {
309            server_url: "http://127.0.0.1:8765".parse().unwrap(),
310            data_dir: tmp.path().to_path_buf(),
311            segments: 0,
312        };
313
314        match sync(&cfg).await {
315            Err(Error::InvalidConfig(_)) => {}
316            Err(e) => panic!("expected InvalidConfig, got {e}"),
317            Ok(_) => panic!("expected error for zero segments"),
318        }
319    }
320}