1use std::borrow::Cow;
11
12use mkit_core::hash::Hash;
13use mkit_core::refs::RefWriteCondition;
14pub use mkit_core::refs::validate_ref_prefix;
15
16use crate::error::{Code, ServerError};
17
18pub const MAX_REF_NAME_BYTES: usize = mkit_core::refs::MAX_REF_NAME_BYTES;
24
25pub const REF_NAME_TOO_LONG: &str = "ref name too long";
28
29pub const SERVED_REFS_PREFIX: &str = "refs/";
33
34pub const REF_NAME_OUTSIDE_REFS: &str = "ref name must start with refs/ (refs outside refs/ \
40 written by older servers are no longer served; see the migration notes)";
41
42#[must_use]
52pub fn is_served_ref_name(name: &str) -> bool {
53 name.starts_with(SERVED_REFS_PREFIX) && validate_ref_name(name)
54}
55
56#[must_use]
59pub fn validate_ref_name(name: &str) -> bool {
60 mkit_core::refs::validate_ref_name(name)
61}
62
63#[derive(Debug, Clone, Copy, PartialEq, Eq)]
67pub enum RefExpectationWire {
68 Unspecified = 0,
70 Any = 1,
72 Missing = 2,
74 Match = 3,
76}
77
78impl RefExpectationWire {
79 #[must_use]
82 pub const fn from_wire(n: i32) -> Self {
83 match n {
84 1 => Self::Any,
85 2 => Self::Missing,
86 3 => Self::Match,
87 _ => Self::Unspecified,
88 }
89 }
90}
91
92#[derive(Debug, Clone, Copy, PartialEq, Eq)]
94#[non_exhaustive]
95pub enum ConflictReason {
96 Exists,
98 Missing,
100 Mismatch,
102}
103
104#[derive(Debug, Clone, PartialEq, Eq)]
108#[non_exhaustive]
109pub enum CasDecision {
110 Committed,
112 Conflict(ConflictReason),
114 Invalid(&'static str),
116}
117
118#[must_use]
126pub fn evaluate_cas(
127 current: Option<&[u8]>,
128 expectation: RefExpectationWire,
129 expected: Option<&[u8]>,
130) -> CasDecision {
131 match expectation {
132 RefExpectationWire::Any => {
133 if expected.is_some() {
134 return CasDecision::Invalid("expected_id must be empty for ANY");
135 }
136 CasDecision::Committed
137 }
138 RefExpectationWire::Missing => {
139 if expected.is_some() {
140 return CasDecision::Invalid("expected_id must be empty for MISSING");
141 }
142 match current {
143 None => CasDecision::Committed,
144 Some(_) => CasDecision::Conflict(ConflictReason::Exists),
145 }
146 }
147 RefExpectationWire::Match => {
148 let Some(expected) = expected else {
149 return CasDecision::Invalid("expected_id required for MATCH");
150 };
151 match current {
152 None => CasDecision::Conflict(ConflictReason::Missing),
153 Some(cur) if cur != expected => CasDecision::Conflict(ConflictReason::Mismatch),
154 Some(_) => CasDecision::Committed,
155 }
156 }
157 RefExpectationWire::Unspecified => {
158 CasDecision::Invalid("expectation is UNSPECIFIED (protocol error)")
159 }
160 }
161}
162
163#[must_use]
167pub fn evaluate_condition(current: Option<&Hash>, condition: &RefWriteCondition) -> CasDecision {
168 match (condition, current) {
169 (RefWriteCondition::Missing, Some(_)) => CasDecision::Conflict(ConflictReason::Exists),
170 (RefWriteCondition::Match(_), None) => CasDecision::Conflict(ConflictReason::Missing),
171 (RefWriteCondition::Match(want), Some(cur)) if cur != want => {
172 CasDecision::Conflict(ConflictReason::Mismatch)
173 }
174 (RefWriteCondition::Any | RefWriteCondition::Missing | RefWriteCondition::Match(_), _) => {
175 CasDecision::Committed
176 }
177 }
178}
179
180#[derive(Debug, Clone, Copy, PartialEq, Eq)]
182#[non_exhaustive]
183pub enum DigestField {
184 PackId,
186 NewId,
188 ExpectedId,
190}
191
192#[derive(Debug, Clone, Copy, PartialEq, Eq)]
194pub enum UnusedExpectedId {
195 Reject,
197 Ignore,
200}
201
202#[derive(Debug, Clone, Copy, PartialEq, Eq)]
206#[non_exhaustive]
207pub enum RefWireError {
208 Unspecified,
210 IdNotEmpty(RefExpectationWire),
212 BadDigest {
214 field: DigestField,
216 len: Option<usize>,
218 },
219}
220
221impl RefWireError {
222 #[must_use]
224 pub const fn code(self) -> Code {
225 Code::InvalidArgument
226 }
227
228 #[must_use]
232 pub const fn ssh_message(self) -> &'static str {
233 match self {
234 Self::Unspecified => "UpdateRef.expectation is required",
235 Self::IdNotEmpty(RefExpectationWire::Missing) => {
236 "expected_id must be empty for MISSING"
237 }
238 Self::IdNotEmpty(_) => "expected_id must be empty for ANY",
239 Self::BadDigest {
240 field: DigestField::PackId,
241 len: None,
242 } => "pack_id missing",
243 Self::BadDigest {
244 field: DigestField::PackId,
245 len: Some(_),
246 } => "pack_id must be 32 bytes",
247 Self::BadDigest {
248 field: DigestField::NewId,
249 ..
250 } => "new_id must be 32 bytes",
251 Self::BadDigest {
252 field: DigestField::ExpectedId,
253 ..
254 } => "MATCH expectation requires a 32-byte expected_id",
255 }
256 }
257
258 #[must_use]
260 pub fn connect_message(self) -> Cow<'static, str> {
261 match self {
262 Self::Unspecified => "expectation MUST NOT be REF_EXPECTATION_UNSPECIFIED".into(),
263 Self::IdNotEmpty(RefExpectationWire::Missing) => {
264 "REF_EXPECTATION_MISSING MUST carry an empty expected_id".into()
265 }
266 Self::IdNotEmpty(_) => "REF_EXPECTATION_ANY MUST carry an empty expected_id".into(),
267 Self::BadDigest { len, .. } => format!(
268 "expected a 32-byte digest, got {} bytes",
269 len.unwrap_or_default()
270 )
271 .into(),
272 }
273 }
274}
275
276impl From<RefWireError> for ServerError {
277 fn from(err: RefWireError) -> Self {
279 Self::new(err.code(), err.connect_message())
280 }
281}
282
283pub fn condition_from_wire(
293 expectation: i32,
294 expected_id: &[u8],
295 unused: UnusedExpectedId,
296) -> Result<RefWriteCondition, RefWireError> {
297 let expectation = RefExpectationWire::from_wire(expectation);
298 let reject_id = unused == UnusedExpectedId::Reject && !expected_id.is_empty();
299 match expectation {
300 RefExpectationWire::Any | RefExpectationWire::Missing if reject_id => {
301 Err(RefWireError::IdNotEmpty(expectation))
302 }
303 RefExpectationWire::Any => Ok(RefWriteCondition::Any),
304 RefExpectationWire::Missing => Ok(RefWriteCondition::Missing),
305 RefExpectationWire::Match => Ok(RefWriteCondition::Match(hash_from_slice(
306 DigestField::ExpectedId,
307 Some(expected_id),
308 )?)),
309 RefExpectationWire::Unspecified => Err(RefWireError::Unspecified),
310 }
311}
312
313pub fn hash_from_slice(field: DigestField, bytes: Option<&[u8]>) -> Result<Hash, RefWireError> {
319 let bytes = bytes.ok_or(RefWireError::BadDigest { field, len: None })?;
320 Hash::try_from(bytes).map_err(|_| RefWireError::BadDigest {
321 field,
322 len: Some(bytes.len()),
323 })
324}
325
326#[must_use]
334pub fn list_scan_prefix(prefix: &str) -> String {
335 let trimmed = prefix.trim_end_matches('/');
336 if trimmed.is_empty() {
337 String::new()
338 } else {
339 format!("{trimmed}/")
340 }
341}
342
343#[must_use]
355pub fn strip_listed_prefix<'a>(full: &'a str, prefix: &str) -> Option<&'a str> {
356 full.strip_prefix(list_scan_prefix(prefix).as_str())
357}
358
359#[cfg(test)]
360mod tests {
361 use super::*;
362
363 const ID_A: &[u8] = &[0xaa; 32];
364 const ID_B: &[u8] = &[0xbb; 32];
365
366 #[test]
368 fn any_clobbers() {
369 assert_eq!(
370 evaluate_cas(Some(ID_A), RefExpectationWire::Any, None),
371 CasDecision::Committed
372 );
373 assert_eq!(
374 evaluate_cas(None, RefExpectationWire::Any, None),
375 CasDecision::Committed
376 );
377 assert!(matches!(
378 evaluate_cas(Some(ID_A), RefExpectationWire::Any, Some(ID_A)),
379 CasDecision::Invalid(_)
380 ));
381 }
382
383 #[test]
385 fn missing_create_only() {
386 assert_eq!(
387 evaluate_cas(None, RefExpectationWire::Missing, None),
388 CasDecision::Committed
389 );
390 assert_eq!(
391 evaluate_cas(Some(ID_A), RefExpectationWire::Missing, None),
392 CasDecision::Conflict(ConflictReason::Exists)
393 );
394 assert!(matches!(
395 evaluate_cas(None, RefExpectationWire::Missing, Some(ID_A)),
396 CasDecision::Invalid(_)
397 ));
398 }
399
400 #[test]
402 fn match_cas() {
403 assert_eq!(
404 evaluate_cas(Some(ID_A), RefExpectationWire::Match, Some(ID_A)),
405 CasDecision::Committed
406 );
407 assert_eq!(
408 evaluate_cas(Some(ID_B), RefExpectationWire::Match, Some(ID_A)),
409 CasDecision::Conflict(ConflictReason::Mismatch)
410 );
411 assert_eq!(
412 evaluate_cas(None, RefExpectationWire::Match, Some(ID_A)),
413 CasDecision::Conflict(ConflictReason::Missing)
414 );
415 assert!(matches!(
416 evaluate_cas(Some(ID_A), RefExpectationWire::Match, None),
417 CasDecision::Invalid(_)
418 ));
419 }
420
421 #[test]
423 fn unspecified_is_protocol_error() {
424 assert!(matches!(
425 evaluate_cas(None, RefExpectationWire::Unspecified, None),
426 CasDecision::Invalid(_)
427 ));
428 }
429
430 #[test]
431 fn evaluate_condition_matches_evaluate_cas() {
432 let (a, b) = ([0xaa; 32], [0xbb; 32]);
433 for current in [None, Some(&a), Some(&b)] {
434 for condition in [
435 RefWriteCondition::Any,
436 RefWriteCondition::Missing,
437 RefWriteCondition::Match(a),
438 RefWriteCondition::Match(b),
439 ] {
440 let (expectation, expected) = match &condition {
441 RefWriteCondition::Any => (RefExpectationWire::Any, None),
442 RefWriteCondition::Missing => (RefExpectationWire::Missing, None),
443 RefWriteCondition::Match(h) => (RefExpectationWire::Match, Some(&h[..])),
444 };
445 assert_eq!(
446 evaluate_condition(current, &condition),
447 evaluate_cas(current.map(|h| &h[..]), expectation, expected),
448 "{current:?} {condition:?}"
449 );
450 }
451 }
452 }
453
454 #[test]
456 fn from_wire_numbers_match_proto() {
457 assert_eq!(RefExpectationWire::from_wire(1), RefExpectationWire::Any);
458 assert_eq!(
459 RefExpectationWire::from_wire(2),
460 RefExpectationWire::Missing
461 );
462 assert_eq!(RefExpectationWire::from_wire(3), RefExpectationWire::Match);
463 assert_eq!(
464 RefExpectationWire::from_wire(0),
465 RefExpectationWire::Unspecified
466 );
467 assert_eq!(
468 RefExpectationWire::from_wire(99),
469 RefExpectationWire::Unspecified
470 );
471 }
472
473 #[test]
477 fn digest_length() {
478 let new_id = |b: &[u8]| hash_from_slice(DigestField::NewId, Some(b));
479 assert_eq!(new_id(&[7; 32]).unwrap(), [7; 32]);
480 for len in [0, 31, 33] {
481 let err = new_id(&vec![0; len]).unwrap_err();
482 assert_eq!(err.code(), Code::InvalidArgument);
483 assert_eq!(
484 err.connect_message(),
485 format!("expected a 32-byte digest, got {len} bytes")
486 );
487 assert_eq!(err.ssh_message(), "new_id must be 32 bytes");
488 }
489 }
490
491 #[test]
495 fn pack_key_from_id_rejects_bad_length_as_invalid_request() {
496 let pack_id = |b: Option<&[u8]>| hash_from_slice(DigestField::PackId, b);
497 let err = pack_id(Some(&[0; 16])).unwrap_err();
498 assert_eq!(err.code(), Code::InvalidArgument);
499 assert_eq!(err.ssh_message(), "pack_id must be 32 bytes");
500 assert_eq!(pack_id(None).unwrap_err().ssh_message(), "pack_id missing");
501 assert_eq!(pack_id(Some(&[7; 32])).unwrap(), [7; 32]);
502 }
503
504 fn rejected(expectation: i32, expected_id: &[u8]) -> RefWireError {
505 condition_from_wire(expectation, expected_id, UnusedExpectedId::Reject).unwrap_err()
506 }
507
508 #[test]
512 fn condition_from_wire_unspecified_and_unknown() {
513 for expectation in [0, 99, -1] {
514 let err = rejected(expectation, &[]);
515 assert_eq!(err, RefWireError::Unspecified);
516 assert_eq!(
517 err.connect_message(),
518 "expectation MUST NOT be REF_EXPECTATION_UNSPECIFIED"
519 );
520 assert_eq!(err.ssh_message(), "UpdateRef.expectation is required");
521 }
522 }
523
524 #[test]
525 fn condition_from_wire_any_or_missing_with_an_id() {
526 assert_eq!(
527 rejected(1, ID_A).connect_message(),
528 "REF_EXPECTATION_ANY MUST carry an empty expected_id"
529 );
530 assert_eq!(
531 rejected(2, ID_A).connect_message(),
532 "REF_EXPECTATION_MISSING MUST carry an empty expected_id"
533 );
534 let err: ServerError = rejected(1, ID_A).into();
535 assert_eq!(err.code(), Code::InvalidArgument);
536 }
537
538 #[test]
539 fn condition_from_ssh_wire_ignores_an_unused_id() {
540 let ssh = |e, id| condition_from_wire(e, id, UnusedExpectedId::Ignore);
541 assert_eq!(ssh(1, ID_A), Ok(RefWriteCondition::Any));
542 assert_eq!(ssh(2, &[1, 2, 3]), Ok(RefWriteCondition::Missing));
543 assert_eq!(ssh(0, &[]), Err(RefWireError::Unspecified));
544 assert_eq!(
545 ssh(3, &[0; 31]).unwrap_err().ssh_message(),
546 "MATCH expectation requires a 32-byte expected_id"
547 );
548 }
549
550 #[test]
551 fn condition_from_wire_match_needs_32_bytes() {
552 assert_eq!(
553 rejected(3, &[0; 31]).connect_message(),
554 "expected a 32-byte digest, got 31 bytes"
555 );
556 assert_eq!(
557 rejected(3, &[]).connect_message(),
558 "expected a 32-byte digest, got 0 bytes"
559 );
560 }
561
562 #[test]
563 fn condition_from_wire_ok() {
564 let connect = |e, id| condition_from_wire(e, id, UnusedExpectedId::Reject);
565 assert_eq!(connect(1, &[]), Ok(RefWriteCondition::Any));
566 assert_eq!(connect(2, &[]), Ok(RefWriteCondition::Missing));
567 assert_eq!(connect(3, ID_B), Ok(RefWriteCondition::Match([0xbb; 32])));
568 }
569
570 #[test]
573 fn strip_listed_prefix_matches_at_component_boundaries() {
574 let strip = strip_listed_prefix;
575 for p in ["refs/heads", "refs/heads/", "refs/heads//"] {
576 assert_eq!(strip("refs/heads/main", p), Some("main"), "{p}");
577 assert_eq!(strip("refs/heads/feat/x", p), Some("feat/x"), "{p}");
578 assert_eq!(strip("refs/tags/v1", p), None, "{p}");
579 }
580 assert_eq!(strip("refs/heads/main", ""), Some("refs/heads/main"));
581 assert_eq!(strip("refs/heads/main", "refs//"), Some("heads/main"));
582 assert_eq!(strip("refs/heads/main", "refs"), Some("heads/main"));
583 assert_eq!(strip("refs/heads/main", "refs/heads/ma"), None);
586 assert_eq!(strip("refs/heads/featx", "refs/heads/feat"), None);
587 assert_eq!(strip("refs/heads/feat/x", "refs/heads/feat"), Some("x"));
588 assert_eq!(strip("refs/heads/main", "refs/heads/main"), None);
590 assert_eq!(list_scan_prefix(""), "");
591 assert_eq!(list_scan_prefix("refs//"), "refs/");
592 assert_eq!(list_scan_prefix("refs/heads/main"), "refs/heads/main/");
593 }
594
595 #[test]
596 fn ref_name_validation_is_mkit_core() {
597 assert!(validate_ref_name("refs/heads/main"));
598 assert!(!validate_ref_name("refs/heads/../main"));
599 let longest = format!("refs/heads/{}", "a".repeat(MAX_REF_NAME_BYTES - 11));
600 assert!(validate_ref_name(&longest));
601 assert!(!validate_ref_name(&format!("{longest}a")));
602 assert!(validate_ref_prefix(""));
603 assert!(validate_ref_prefix("refs/heads/"));
604 assert!(!validate_ref_prefix("/"));
605 }
606}