1pub mod topology;
8
9use std::collections::HashMap;
10
11#[derive(Clone, Debug, Eq, PartialEq)]
12pub struct CoordinatorClaim {
13 pub model_id: String,
14 pub package_ref: String,
15 pub manifest_sha256: String,
16 pub topology_id: String,
17 pub run_id: String,
18 pub coordinator_id: String,
19 pub coordinator_term: u64,
20 pub participant_set_hash: String,
21 pub topology_hash: String,
22 pub lease_until_unix_ms: u64,
23}
24
25#[derive(Clone, Debug, Eq, PartialEq)]
26pub struct LoadClaimRef {
27 pub model_id: String,
28 pub package_ref: String,
29 pub manifest_sha256: String,
30 pub topology_id: String,
31 pub run_id: String,
32 pub coordinator_id: Option<String>,
33 pub coordinator_term: u64,
34}
35
36#[derive(Clone, Debug, Eq, PartialEq)]
37pub enum ClaimDecision {
38 Accepted {
39 supersedes_term: Option<u64>,
40 claim: CoordinatorClaim,
41 },
42 Rejected {
43 current: Option<CoordinatorClaim>,
44 reason: ClaimRejection,
45 },
46}
47
48#[derive(Clone, Debug, Eq, PartialEq, thiserror::Error)]
49pub enum ClaimRejection {
50 #[error("coordinator claim requires model_id")]
51 MissingModelId,
52 #[error("coordinator claim requires package_ref")]
53 MissingPackageRef,
54 #[error("coordinator claim requires manifest_sha256")]
55 MissingManifestSha256,
56 #[error("coordinator claim requires topology_id and run_id")]
57 MissingTopologyRun,
58 #[error("coordinator claim requires coordinator_id")]
59 MissingCoordinatorId,
60 #[error("coordinator claim requires non-zero term")]
61 MissingTerm,
62 #[error("coordinator claim requires participant and topology hashes")]
63 MissingHashes,
64 #[error("coordinator claim lease is expired")]
65 ExpiredLease,
66 #[error("stale coordinator term {claim_term} < {current_term}")]
67 StaleTerm { claim_term: u64, current_term: u64 },
68 #[error("conflicting coordinator claim for existing term")]
69 ConflictingSameTerm,
70}
71
72#[derive(Clone, Debug, Eq, PartialEq, thiserror::Error)]
73pub enum LoadRejection {
74 #[error("missing coordinator term")]
75 MissingTerm,
76 #[error("missing coordinator id")]
77 MissingCoordinatorId,
78 #[error("missing coordinator claim")]
79 MissingClaim,
80 #[error("coordinator term mismatch: load={load_term} claim={claim_term}")]
81 TermMismatch { load_term: u64, claim_term: u64 },
82 #[error("coordinator id mismatch")]
83 CoordinatorMismatch,
84 #[error("coordinator claim does not match topology/run")]
85 TopologyRunMismatch,
86 #[error("coordinator lease expired")]
87 ExpiredLease,
88}
89
90#[derive(Clone, Debug, Default)]
91pub struct ClaimFence {
92 claims: HashMap<ClaimKey, CoordinatorClaim>,
93}
94
95impl ClaimFence {
96 pub fn accept_claim(&mut self, claim: CoordinatorClaim, now_unix_ms: u64) -> ClaimDecision {
97 if let Some(reason) = validate_claim_shape(&claim, now_unix_ms) {
98 return ClaimDecision::Rejected {
99 current: None,
100 reason,
101 };
102 }
103
104 let key = ClaimKey::from_claim(&claim);
105 if let Some(current) = self.claims.get(&key) {
106 if claim.coordinator_term < current.coordinator_term {
107 return ClaimDecision::Rejected {
108 current: Some(current.clone()),
109 reason: ClaimRejection::StaleTerm {
110 claim_term: claim.coordinator_term,
111 current_term: current.coordinator_term,
112 },
113 };
114 }
115 if claim.coordinator_term == current.coordinator_term
116 && !same_claim_epoch(&claim, current)
117 {
118 return ClaimDecision::Rejected {
119 current: Some(current.clone()),
120 reason: ClaimRejection::ConflictingSameTerm,
121 };
122 }
123 }
124
125 let previous = self.claims.insert(key, claim.clone());
126 ClaimDecision::Accepted {
127 supersedes_term: previous.as_ref().and_then(|current| {
128 (claim.coordinator_term > current.coordinator_term)
129 .then_some(current.coordinator_term)
130 }),
131 claim,
132 }
133 }
134
135 pub fn validate_load(
136 &self,
137 load: &LoadClaimRef,
138 now_unix_ms: u64,
139 ) -> Result<(), LoadRejection> {
140 if load.coordinator_term == 0 {
141 return Err(LoadRejection::MissingTerm);
142 }
143 let Some(coordinator_id) = load.coordinator_id.as_deref() else {
144 return Err(LoadRejection::MissingCoordinatorId);
145 };
146 let key = ClaimKey::from_load(load);
147 let Some(claim) = self.claims.get(&key) else {
148 return Err(LoadRejection::MissingClaim);
149 };
150 if claim.coordinator_term != load.coordinator_term {
151 return Err(LoadRejection::TermMismatch {
152 load_term: load.coordinator_term,
153 claim_term: claim.coordinator_term,
154 });
155 }
156 if claim.coordinator_id != coordinator_id {
157 return Err(LoadRejection::CoordinatorMismatch);
158 }
159 if claim.topology_id != load.topology_id || claim.run_id != load.run_id {
160 return Err(LoadRejection::TopologyRunMismatch);
161 }
162 if claim.lease_until_unix_ms < now_unix_ms {
163 return Err(LoadRejection::ExpiredLease);
164 }
165 Ok(())
166 }
167
168 pub fn current_claim_for(
169 &self,
170 model_id: &str,
171 package_ref: &str,
172 manifest_sha256: &str,
173 ) -> Option<&CoordinatorClaim> {
174 self.claims.get(&ClaimKey {
175 model_id: model_id.to_string(),
176 package_ref: package_ref.to_string(),
177 manifest_sha256: manifest_sha256.to_string(),
178 })
179 }
180}
181
182pub fn quorum_requirement(planned_stage_count: usize) -> usize {
183 planned_stage_count / 2 + 1
184}
185
186pub fn same_claim_epoch(left: &CoordinatorClaim, right: &CoordinatorClaim) -> bool {
187 left.model_id == right.model_id
188 && left.package_ref == right.package_ref
189 && left.manifest_sha256 == right.manifest_sha256
190 && left.topology_id == right.topology_id
191 && left.run_id == right.run_id
192 && left.coordinator_id == right.coordinator_id
193 && left.participant_set_hash == right.participant_set_hash
194 && left.topology_hash == right.topology_hash
195}
196
197fn validate_claim_shape(claim: &CoordinatorClaim, now_unix_ms: u64) -> Option<ClaimRejection> {
198 if claim.model_id.is_empty() {
199 return Some(ClaimRejection::MissingModelId);
200 }
201 if claim.package_ref.is_empty() {
202 return Some(ClaimRejection::MissingPackageRef);
203 }
204 if claim.manifest_sha256.is_empty() {
205 return Some(ClaimRejection::MissingManifestSha256);
206 }
207 if claim.topology_id.is_empty() || claim.run_id.is_empty() {
208 return Some(ClaimRejection::MissingTopologyRun);
209 }
210 if claim.coordinator_id.is_empty() {
211 return Some(ClaimRejection::MissingCoordinatorId);
212 }
213 if claim.coordinator_term == 0 {
214 return Some(ClaimRejection::MissingTerm);
215 }
216 if claim.participant_set_hash.is_empty() || claim.topology_hash.is_empty() {
217 return Some(ClaimRejection::MissingHashes);
218 }
219 if claim.lease_until_unix_ms <= now_unix_ms {
220 return Some(ClaimRejection::ExpiredLease);
221 }
222 None
223}
224
225#[derive(Clone, Debug, Eq, Hash, PartialEq)]
226struct ClaimKey {
227 model_id: String,
228 package_ref: String,
229 manifest_sha256: String,
230}
231
232impl ClaimKey {
233 fn from_claim(claim: &CoordinatorClaim) -> Self {
234 Self {
235 model_id: claim.model_id.clone(),
236 package_ref: claim.package_ref.clone(),
237 manifest_sha256: claim.manifest_sha256.clone(),
238 }
239 }
240
241 fn from_load(load: &LoadClaimRef) -> Self {
242 Self {
243 model_id: load.model_id.clone(),
244 package_ref: load.package_ref.clone(),
245 manifest_sha256: load.manifest_sha256.clone(),
246 }
247 }
248}
249
250#[cfg(test)]
251mod tests {
252 use super::*;
253
254 fn claim(term: u64) -> CoordinatorClaim {
255 CoordinatorClaim {
256 model_id: "model".to_string(),
257 package_ref: "hf://pkg".to_string(),
258 manifest_sha256: "manifest".to_string(),
259 topology_id: format!("topology-{term}"),
260 run_id: format!("run-{term}"),
261 coordinator_id: "node-a".to_string(),
262 coordinator_term: term,
263 participant_set_hash: "participants".to_string(),
264 topology_hash: format!("topology-hash-{term}"),
265 lease_until_unix_ms: 10_000,
266 }
267 }
268
269 fn load(term: u64) -> LoadClaimRef {
270 LoadClaimRef {
271 model_id: "model".to_string(),
272 package_ref: "hf://pkg".to_string(),
273 manifest_sha256: "manifest".to_string(),
274 topology_id: format!("topology-{term}"),
275 run_id: format!("run-{term}"),
276 coordinator_id: Some("node-a".to_string()),
277 coordinator_term: term,
278 }
279 }
280
281 #[test]
282 fn accepts_first_valid_claim() {
283 let mut fence = ClaimFence::default();
284 assert!(matches!(
285 fence.accept_claim(claim(1), 1_000),
286 ClaimDecision::Accepted {
287 supersedes_term: None,
288 ..
289 }
290 ));
291 }
292
293 #[test]
294 fn rejects_stale_claim_after_newer_term() {
295 let mut fence = ClaimFence::default();
296 fence.accept_claim(claim(2), 1_000);
297 assert!(matches!(
298 fence.accept_claim(claim(1), 1_000),
299 ClaimDecision::Rejected {
300 reason: ClaimRejection::StaleTerm {
301 claim_term: 1,
302 current_term: 2
303 },
304 ..
305 }
306 ));
307 }
308
309 #[test]
310 fn rejects_conflicting_claim_for_same_term() {
311 let mut fence = ClaimFence::default();
312 fence.accept_claim(claim(1), 1_000);
313 let mut conflicting = claim(1);
314 conflicting.coordinator_id = "node-b".to_string();
315 assert!(matches!(
316 fence.accept_claim(conflicting, 1_000),
317 ClaimDecision::Rejected {
318 reason: ClaimRejection::ConflictingSameTerm,
319 ..
320 }
321 ));
322 }
323
324 #[test]
325 fn newer_claim_supersedes_old_term() {
326 let mut fence = ClaimFence::default();
327 fence.accept_claim(claim(1), 1_000);
328 assert!(matches!(
329 fence.accept_claim(claim(2), 1_000),
330 ClaimDecision::Accepted {
331 supersedes_term: Some(1),
332 ..
333 }
334 ));
335 }
336
337 #[test]
338 fn validates_load_against_current_claim() {
339 let mut fence = ClaimFence::default();
340 fence.accept_claim(claim(3), 1_000);
341 assert_eq!(fence.validate_load(&load(3), 1_000), Ok(()));
342 }
343
344 #[test]
345 fn rejects_load_without_matching_claim() {
346 let fence = ClaimFence::default();
347 assert_eq!(
348 fence.validate_load(&load(1), 1_000),
349 Err(LoadRejection::MissingClaim)
350 );
351 }
352
353 #[test]
354 fn rejects_load_for_stale_term() {
355 let mut fence = ClaimFence::default();
356 fence.accept_claim(claim(2), 1_000);
357 assert_eq!(
358 fence.validate_load(&load(1), 1_000),
359 Err(LoadRejection::TermMismatch {
360 load_term: 1,
361 claim_term: 2,
362 })
363 );
364 }
365
366 #[test]
367 fn rejects_expired_claim_and_expired_load() {
368 let mut fence = ClaimFence::default();
369 assert!(matches!(
370 fence.accept_claim(claim(1), 10_001),
371 ClaimDecision::Rejected {
372 reason: ClaimRejection::ExpiredLease,
373 ..
374 }
375 ));
376
377 fence.accept_claim(claim(2), 1_000);
378 assert_eq!(
379 fence.validate_load(&load(2), 10_001),
380 Err(LoadRejection::ExpiredLease)
381 );
382 }
383
384 #[test]
385 fn quorum_is_majority_of_planned_stages() {
386 assert_eq!(quorum_requirement(1), 1);
387 assert_eq!(quorum_requirement(2), 2);
388 assert_eq!(quorum_requirement(3), 2);
389 assert_eq!(quorum_requirement(4), 3);
390 }
391}