1use crate::{SnapAttemptStore, SnapGeneration, SnapPivotPolicy, SnapSyncError, SnapWrite};
4use reth_storage_api::{BlockHashReader, HeaderProvider, MetadataProvider, MetadataWriter};
5use tokio_util::sync::CancellationToken;
6
7#[derive(Debug)]
12pub struct SnapSyncSession {
13 policy: SnapPivotPolicy,
15 state: SnapSyncSessionState,
17 cancellation: CancellationToken,
19}
20
21impl SnapSyncSession {
22 pub fn new(policy: SnapPivotPolicy) -> Self {
24 Self {
25 policy,
26 state: SnapSyncSessionState::Waiting,
27 cancellation: CancellationToken::new(),
28 }
29 }
30
31 pub const fn state(&self) -> &SnapSyncSessionState {
33 &self.state
34 }
35
36 pub const fn target(&self) -> Option<&SnapGeneration> {
38 self.state.target()
39 }
40
41 pub fn is_cancelled(&self) -> bool {
43 self.cancellation.is_cancelled()
44 }
45
46 pub fn select(
52 &mut self,
53 provider: &impl HeaderProvider,
54 head: u64,
55 finalized: Option<u64>,
56 ) -> Result<&SnapSyncSessionState, SnapSyncError> {
57 if matches!(self.state, SnapSyncSessionState::Waiting | SnapSyncSessionState::Selected(_)) {
58 self.state = match self.policy.select(provider, head, finalized)? {
59 Some(generation) => SnapSyncSessionState::Selected(generation),
60 None => SnapSyncSessionState::Waiting,
61 };
62 }
63 Ok(&self.state)
64 }
65
66 pub fn start(&mut self) -> Option<(SnapGeneration, CancellationToken)> {
71 let SnapSyncSessionState::Selected(generation) = self.state else { return None };
72 self.state = SnapSyncSessionState::Downloading(generation);
73 Some((generation, self.cancellation.clone()))
74 }
75
76 pub fn resume(&mut self, generation: SnapGeneration) -> Option<CancellationToken> {
81 if !matches!(self.state, SnapSyncSessionState::Waiting | SnapSyncSessionState::Selected(_))
82 {
83 return None
84 }
85 self.state = SnapSyncSessionState::Downloading(generation);
86 Some(self.cancellation.clone())
87 }
88
89 pub fn advance<P>(
95 &mut self,
96 provider: &P,
97 write: SnapWrite,
98 head: u64,
99 finalized: Option<u64>,
100 ) -> Result<Option<SnapWrite>, SnapSyncError>
101 where
102 P: HeaderProvider + BlockHashReader + MetadataProvider + MetadataWriter,
103 {
104 let SnapSyncSessionState::Downloading(mut current) = self.state else { return Ok(None) };
105 let attempt = provider.authorize_snap_write(write)?;
106 if attempt.pivot() != current.target() {
107 current = SnapGeneration::new(attempt.pivot(), attempt.state_root());
108 self.state = SnapSyncSessionState::Downloading(current);
109 }
110 if !self.policy.needs_advance(current, head) {
111 return Ok(None)
112 }
113 let Some(next) = self.policy.select(provider, head, finalized)? else { return Ok(None) };
114 if next.target().number <= current.target().number {
116 return Ok(None)
117 }
118 let write = provider.advance_snap_pivot(write, next)?;
119 self.state = SnapSyncSessionState::Downloading(next);
120 Ok(Some(write))
121 }
122
123 pub fn cancel(&mut self) {
127 self.cancellation.cancel();
128 self.state = SnapSyncSessionState::Cancelled;
129 }
130}
131
132#[derive(Clone, Copy, Debug, Eq, PartialEq)]
134pub enum SnapSyncSessionState {
135 Waiting,
137 Selected(SnapGeneration),
139 Downloading(SnapGeneration),
141 Cancelled,
143}
144
145impl SnapSyncSessionState {
146 pub const fn target(&self) -> Option<&SnapGeneration> {
148 match self {
149 Self::Selected(generation) | Self::Downloading(generation) => Some(generation),
150 Self::Waiting | Self::Cancelled => None,
151 }
152 }
153}
154
155#[cfg(test)]
156mod tests {
157 use super::*;
158 use crate::test_utils::{chain, hashed_factory, policy, provider_with};
159 use reth_primitives_traits::SealedHeader;
160 use reth_provider::{
161 test_utils::{insert_headers, MockNodeTypesWithDB},
162 DBProvider, DatabaseProviderFactory, ProviderFactory,
163 };
164
165 fn session() -> SnapSyncSession {
166 SnapSyncSession::new(policy())
167 }
168
169 fn downloading() -> (ProviderFactory<MockNodeTypesWithDB>, SnapSyncSession, SnapWrite) {
172 let factory = hashed_factory();
173 let headers: Vec<_> = chain(Some(0)).into_iter().map(SealedHeader::seal_slow).collect();
174 insert_headers(&factory, &headers);
175 let provider = factory.database_provider_rw().unwrap();
176 let mut session = SnapSyncSession::new(policy().with_advance_after(1));
177 session.select(&provider, 2, None).unwrap();
178 let (generation, _) = session.start().unwrap();
179 let write = provider.start_snap_attempt(generation).unwrap();
180 provider.commit().unwrap();
181 (factory, session, write)
182 }
183
184 #[test]
185 fn waits_while_no_block_is_eligible() {
186 let provider = provider_with(chain(Some(3)));
188 let mut session = session();
189
190 assert_eq!(session.select(&provider, 3, None).unwrap(), &SnapSyncSessionState::Waiting);
191 assert_eq!(session.target(), None);
192 assert!(session.start().is_none());
193 }
194
195 #[test]
196 fn selects_an_eligible_target() {
197 let headers = chain(Some(0));
198 let expected = headers[2].clone();
199 let provider = provider_with(headers);
200 let mut session = session();
201
202 session.select(&provider, 3, None).unwrap();
203
204 let target = session.target().unwrap();
205 assert_eq!(target.target().number, 2);
206 assert_eq!(target.target().hash, expected.hash_slow());
207 }
208
209 #[test]
210 fn a_target_no_work_took_is_replaced_as_the_head_advances() {
211 let provider = provider_with(chain(Some(0)));
212 let mut session = session();
213
214 session.select(&provider, 2, None).unwrap();
215 assert_eq!(session.target().unwrap().target().number, 1);
216
217 session.select(&provider, 3, None).unwrap();
218 assert_eq!(session.target().unwrap().target().number, 2);
219 }
220
221 #[test]
222 fn a_target_no_work_took_is_dropped_once_it_is_no_longer_eligible() {
223 let provider = provider_with(chain(Some(0)));
224 let mut session = session();
225 session.select(&provider, 3, None).unwrap();
226
227 session.select(&provider, 9, None).unwrap();
229
230 assert_eq!(session.state(), &SnapSyncSessionState::Waiting);
231 assert!(session.start().is_none());
232 }
233
234 #[test]
235 fn a_target_work_took_is_kept() {
236 let provider = provider_with(chain(Some(0)));
237 let mut session = session();
238 session.select(&provider, 2, None).unwrap();
239 let (started, _) = session.start().unwrap();
240
241 session.select(&provider, 3, None).unwrap();
242
243 assert_eq!(session.state(), &SnapSyncSessionState::Downloading(started));
244 }
245
246 #[test]
247 fn only_one_worker_takes_a_target() {
248 let provider = provider_with(chain(Some(0)));
249 let mut session = session();
250 session.select(&provider, 3, None).unwrap();
251
252 assert!(session.start().is_some());
253 assert!(session.start().is_none());
254 }
255
256 #[test]
257 fn a_recorded_target_is_resumed_once() {
258 let (factory, _, write) = downloading();
259 let provider = factory.database_provider_rw().unwrap();
260 let attempt = provider.authorize_snap_write(write).unwrap();
261 let recorded = SnapGeneration::new(attempt.pivot(), attempt.state_root());
262 let mut session = SnapSyncSession::new(policy().with_advance_after(1));
263
264 assert!(session.resume(recorded).is_some());
265 assert!(session.resume(recorded).is_none());
266 assert!(session.start().is_none());
267 assert_eq!(session.state(), &SnapSyncSessionState::Downloading(recorded));
268 assert!(session.advance(&provider, write, 3, None).unwrap().is_some());
270 assert_eq!(session.target().unwrap().target().number, 2);
271 }
272
273 #[test]
274 fn cancellation_stops_outstanding_work() {
275 let provider = provider_with(chain(Some(0)));
276 let mut session = session();
277 session.select(&provider, 3, None).unwrap();
278 let (_, outstanding) = session.start().unwrap();
279
280 session.cancel();
281
282 assert!(outstanding.is_cancelled());
283 assert!(session.is_cancelled());
284 assert_eq!(session.state(), &SnapSyncSessionState::Cancelled);
285 assert_eq!(session.target(), None);
286 }
287
288 #[test]
289 fn a_cancelled_session_selects_nothing() {
290 let provider = provider_with(chain(Some(0)));
291 let mut session = session();
292 session.cancel();
293
294 assert_eq!(session.select(&provider, 3, None).unwrap(), &SnapSyncSessionState::Cancelled);
295 assert!(session.start().is_none());
296 }
297
298 #[test]
299 fn a_lagging_target_advances_with_the_attempt() {
300 let (factory, mut session, write) = downloading();
301 let provider = factory.database_provider_rw().unwrap();
302
303 let advanced = session.advance(&provider, write, 3, None).unwrap().unwrap();
304
305 let target = *session.target().unwrap();
306 assert_eq!(target.target().number, 2);
307 assert!(matches!(session.state(), SnapSyncSessionState::Downloading(_)));
308 let attempt = provider.authorize_snap_write(advanced).unwrap();
309 assert_eq!((attempt.pivot(), attempt.state_root()), (target.target(), target.state_root()));
310 assert!(matches!(
311 provider.authorize_snap_write(write),
312 Err(SnapSyncError::StaleWrite { .. })
313 ));
314 }
315
316 #[test]
317 fn a_target_within_the_window_is_kept() {
318 let (factory, mut session, write) = downloading();
319 let provider = factory.database_provider_rw().unwrap();
320 let before = *session.target().unwrap();
321
322 assert_eq!(session.advance(&provider, write, 2, None).unwrap(), None);
323
324 assert_eq!(session.target(), Some(&before));
325 provider.authorize_snap_write(write).unwrap();
326 }
327
328 #[test]
329 fn a_refused_advance_keeps_the_target() {
330 let (factory, mut session, write) = downloading();
331 let provider = factory.database_provider_rw().unwrap();
332 let before = *session.target().unwrap();
333 provider.abandon_snap_attempt().unwrap();
334
335 let refused = session.advance(&provider, write, 3, None);
336
337 assert!(matches!(refused, Err(SnapSyncError::StaleWrite { .. })));
338 assert_eq!(session.state(), &SnapSyncSessionState::Downloading(before));
339 }
340
341 #[test]
342 fn a_rolled_back_advance_is_retried() {
343 let (factory, mut session, write) = downloading();
344 let provider = factory.database_provider_rw().unwrap();
345 session.advance(&provider, write, 3, None).unwrap().unwrap();
346 drop(provider);
347
348 let provider = factory.database_provider_rw().unwrap();
349 let advanced = session.advance(&provider, write, 3, None).unwrap().unwrap();
350
351 assert_eq!(session.target().unwrap().target().number, 2);
352 assert_eq!(provider.authorize_snap_write(advanced).unwrap().pivot().number, 2);
353 }
354
355 #[test]
356 fn a_target_no_work_took_is_not_advanced() {
357 let (factory, _, write) = downloading();
358 let provider = factory.database_provider_rw().unwrap();
359 let mut session = SnapSyncSession::new(policy().with_advance_after(1));
360 session.select(&provider, 2, None).unwrap();
361
362 assert_eq!(session.advance(&provider, write, 3, None).unwrap(), None);
363 assert_eq!(session.target().unwrap().target().number, 1);
364 }
365}