use std::borrow::Cow;
use mkit_core::hash::Hash;
use mkit_core::refs::RefWriteCondition;
pub use mkit_core::refs::validate_ref_prefix;
use crate::error::{Code, ServerError};
pub const MAX_REF_NAME_BYTES: usize = mkit_core::refs::MAX_REF_NAME_BYTES;
pub const REF_NAME_TOO_LONG: &str = "ref name too long";
pub const SERVED_REFS_PREFIX: &str = "refs/";
pub const REF_NAME_OUTSIDE_REFS: &str = "ref name must start with refs/ (refs outside refs/ \
written by older servers are no longer served; see the migration notes)";
#[must_use]
pub fn is_served_ref_name(name: &str) -> bool {
name.starts_with(SERVED_REFS_PREFIX) && validate_ref_name(name)
}
#[must_use]
pub fn validate_ref_name(name: &str) -> bool {
mkit_core::refs::validate_ref_name(name)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RefExpectationWire {
Unspecified = 0,
Any = 1,
Missing = 2,
Match = 3,
}
impl RefExpectationWire {
#[must_use]
pub const fn from_wire(n: i32) -> Self {
match n {
1 => Self::Any,
2 => Self::Missing,
3 => Self::Match,
_ => Self::Unspecified,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum ConflictReason {
Exists,
Missing,
Mismatch,
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum CasDecision {
Committed,
Conflict(ConflictReason),
Invalid(&'static str),
}
#[must_use]
pub fn evaluate_cas(
current: Option<&[u8]>,
expectation: RefExpectationWire,
expected: Option<&[u8]>,
) -> CasDecision {
match expectation {
RefExpectationWire::Any => {
if expected.is_some() {
return CasDecision::Invalid("expected_id must be empty for ANY");
}
CasDecision::Committed
}
RefExpectationWire::Missing => {
if expected.is_some() {
return CasDecision::Invalid("expected_id must be empty for MISSING");
}
match current {
None => CasDecision::Committed,
Some(_) => CasDecision::Conflict(ConflictReason::Exists),
}
}
RefExpectationWire::Match => {
let Some(expected) = expected else {
return CasDecision::Invalid("expected_id required for MATCH");
};
match current {
None => CasDecision::Conflict(ConflictReason::Missing),
Some(cur) if cur != expected => CasDecision::Conflict(ConflictReason::Mismatch),
Some(_) => CasDecision::Committed,
}
}
RefExpectationWire::Unspecified => {
CasDecision::Invalid("expectation is UNSPECIFIED (protocol error)")
}
}
}
#[must_use]
pub fn evaluate_condition(current: Option<&Hash>, condition: &RefWriteCondition) -> CasDecision {
match (condition, current) {
(RefWriteCondition::Missing, Some(_)) => CasDecision::Conflict(ConflictReason::Exists),
(RefWriteCondition::Match(_), None) => CasDecision::Conflict(ConflictReason::Missing),
(RefWriteCondition::Match(want), Some(cur)) if cur != want => {
CasDecision::Conflict(ConflictReason::Mismatch)
}
(RefWriteCondition::Any | RefWriteCondition::Missing | RefWriteCondition::Match(_), _) => {
CasDecision::Committed
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum DigestField {
PackId,
NewId,
ExpectedId,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum UnusedExpectedId {
Reject,
Ignore,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum RefWireError {
Unspecified,
IdNotEmpty(RefExpectationWire),
BadDigest {
field: DigestField,
len: Option<usize>,
},
}
impl RefWireError {
#[must_use]
pub const fn code(self) -> Code {
Code::InvalidArgument
}
#[must_use]
pub const fn ssh_message(self) -> &'static str {
match self {
Self::Unspecified => "UpdateRef.expectation is required",
Self::IdNotEmpty(RefExpectationWire::Missing) => {
"expected_id must be empty for MISSING"
}
Self::IdNotEmpty(_) => "expected_id must be empty for ANY",
Self::BadDigest {
field: DigestField::PackId,
len: None,
} => "pack_id missing",
Self::BadDigest {
field: DigestField::PackId,
len: Some(_),
} => "pack_id must be 32 bytes",
Self::BadDigest {
field: DigestField::NewId,
..
} => "new_id must be 32 bytes",
Self::BadDigest {
field: DigestField::ExpectedId,
..
} => "MATCH expectation requires a 32-byte expected_id",
}
}
#[must_use]
pub fn connect_message(self) -> Cow<'static, str> {
match self {
Self::Unspecified => "expectation MUST NOT be REF_EXPECTATION_UNSPECIFIED".into(),
Self::IdNotEmpty(RefExpectationWire::Missing) => {
"REF_EXPECTATION_MISSING MUST carry an empty expected_id".into()
}
Self::IdNotEmpty(_) => "REF_EXPECTATION_ANY MUST carry an empty expected_id".into(),
Self::BadDigest { len, .. } => format!(
"expected a 32-byte digest, got {} bytes",
len.unwrap_or_default()
)
.into(),
}
}
}
impl From<RefWireError> for ServerError {
fn from(err: RefWireError) -> Self {
Self::new(err.code(), err.connect_message())
}
}
pub fn condition_from_wire(
expectation: i32,
expected_id: &[u8],
unused: UnusedExpectedId,
) -> Result<RefWriteCondition, RefWireError> {
let expectation = RefExpectationWire::from_wire(expectation);
let reject_id = unused == UnusedExpectedId::Reject && !expected_id.is_empty();
match expectation {
RefExpectationWire::Any | RefExpectationWire::Missing if reject_id => {
Err(RefWireError::IdNotEmpty(expectation))
}
RefExpectationWire::Any => Ok(RefWriteCondition::Any),
RefExpectationWire::Missing => Ok(RefWriteCondition::Missing),
RefExpectationWire::Match => Ok(RefWriteCondition::Match(hash_from_slice(
DigestField::ExpectedId,
Some(expected_id),
)?)),
RefExpectationWire::Unspecified => Err(RefWireError::Unspecified),
}
}
pub fn hash_from_slice(field: DigestField, bytes: Option<&[u8]>) -> Result<Hash, RefWireError> {
let bytes = bytes.ok_or(RefWireError::BadDigest { field, len: None })?;
Hash::try_from(bytes).map_err(|_| RefWireError::BadDigest {
field,
len: Some(bytes.len()),
})
}
#[must_use]
pub fn list_scan_prefix(prefix: &str) -> String {
let trimmed = prefix.trim_end_matches('/');
if trimmed.is_empty() {
String::new()
} else {
format!("{trimmed}/")
}
}
#[must_use]
pub fn strip_listed_prefix<'a>(full: &'a str, prefix: &str) -> Option<&'a str> {
full.strip_prefix(list_scan_prefix(prefix).as_str())
}
#[cfg(test)]
mod tests {
use super::*;
const ID_A: &[u8] = &[0xaa; 32];
const ID_B: &[u8] = &[0xbb; 32];
#[test]
fn any_clobbers() {
assert_eq!(
evaluate_cas(Some(ID_A), RefExpectationWire::Any, None),
CasDecision::Committed
);
assert_eq!(
evaluate_cas(None, RefExpectationWire::Any, None),
CasDecision::Committed
);
assert!(matches!(
evaluate_cas(Some(ID_A), RefExpectationWire::Any, Some(ID_A)),
CasDecision::Invalid(_)
));
}
#[test]
fn missing_create_only() {
assert_eq!(
evaluate_cas(None, RefExpectationWire::Missing, None),
CasDecision::Committed
);
assert_eq!(
evaluate_cas(Some(ID_A), RefExpectationWire::Missing, None),
CasDecision::Conflict(ConflictReason::Exists)
);
assert!(matches!(
evaluate_cas(None, RefExpectationWire::Missing, Some(ID_A)),
CasDecision::Invalid(_)
));
}
#[test]
fn match_cas() {
assert_eq!(
evaluate_cas(Some(ID_A), RefExpectationWire::Match, Some(ID_A)),
CasDecision::Committed
);
assert_eq!(
evaluate_cas(Some(ID_B), RefExpectationWire::Match, Some(ID_A)),
CasDecision::Conflict(ConflictReason::Mismatch)
);
assert_eq!(
evaluate_cas(None, RefExpectationWire::Match, Some(ID_A)),
CasDecision::Conflict(ConflictReason::Missing)
);
assert!(matches!(
evaluate_cas(Some(ID_A), RefExpectationWire::Match, None),
CasDecision::Invalid(_)
));
}
#[test]
fn unspecified_is_protocol_error() {
assert!(matches!(
evaluate_cas(None, RefExpectationWire::Unspecified, None),
CasDecision::Invalid(_)
));
}
#[test]
fn evaluate_condition_matches_evaluate_cas() {
let (a, b) = ([0xaa; 32], [0xbb; 32]);
for current in [None, Some(&a), Some(&b)] {
for condition in [
RefWriteCondition::Any,
RefWriteCondition::Missing,
RefWriteCondition::Match(a),
RefWriteCondition::Match(b),
] {
let (expectation, expected) = match &condition {
RefWriteCondition::Any => (RefExpectationWire::Any, None),
RefWriteCondition::Missing => (RefExpectationWire::Missing, None),
RefWriteCondition::Match(h) => (RefExpectationWire::Match, Some(&h[..])),
};
assert_eq!(
evaluate_condition(current, &condition),
evaluate_cas(current.map(|h| &h[..]), expectation, expected),
"{current:?} {condition:?}"
);
}
}
}
#[test]
fn from_wire_numbers_match_proto() {
assert_eq!(RefExpectationWire::from_wire(1), RefExpectationWire::Any);
assert_eq!(
RefExpectationWire::from_wire(2),
RefExpectationWire::Missing
);
assert_eq!(RefExpectationWire::from_wire(3), RefExpectationWire::Match);
assert_eq!(
RefExpectationWire::from_wire(0),
RefExpectationWire::Unspecified
);
assert_eq!(
RefExpectationWire::from_wire(99),
RefExpectationWire::Unspecified
);
}
#[test]
fn digest_length() {
let new_id = |b: &[u8]| hash_from_slice(DigestField::NewId, Some(b));
assert_eq!(new_id(&[7; 32]).unwrap(), [7; 32]);
for len in [0, 31, 33] {
let err = new_id(&vec![0; len]).unwrap_err();
assert_eq!(err.code(), Code::InvalidArgument);
assert_eq!(
err.connect_message(),
format!("expected a 32-byte digest, got {len} bytes")
);
assert_eq!(err.ssh_message(), "new_id must be 32 bytes");
}
}
#[test]
fn pack_key_from_id_rejects_bad_length_as_invalid_request() {
let pack_id = |b: Option<&[u8]>| hash_from_slice(DigestField::PackId, b);
let err = pack_id(Some(&[0; 16])).unwrap_err();
assert_eq!(err.code(), Code::InvalidArgument);
assert_eq!(err.ssh_message(), "pack_id must be 32 bytes");
assert_eq!(pack_id(None).unwrap_err().ssh_message(), "pack_id missing");
assert_eq!(pack_id(Some(&[7; 32])).unwrap(), [7; 32]);
}
fn rejected(expectation: i32, expected_id: &[u8]) -> RefWireError {
condition_from_wire(expectation, expected_id, UnusedExpectedId::Reject).unwrap_err()
}
#[test]
fn condition_from_wire_unspecified_and_unknown() {
for expectation in [0, 99, -1] {
let err = rejected(expectation, &[]);
assert_eq!(err, RefWireError::Unspecified);
assert_eq!(
err.connect_message(),
"expectation MUST NOT be REF_EXPECTATION_UNSPECIFIED"
);
assert_eq!(err.ssh_message(), "UpdateRef.expectation is required");
}
}
#[test]
fn condition_from_wire_any_or_missing_with_an_id() {
assert_eq!(
rejected(1, ID_A).connect_message(),
"REF_EXPECTATION_ANY MUST carry an empty expected_id"
);
assert_eq!(
rejected(2, ID_A).connect_message(),
"REF_EXPECTATION_MISSING MUST carry an empty expected_id"
);
let err: ServerError = rejected(1, ID_A).into();
assert_eq!(err.code(), Code::InvalidArgument);
}
#[test]
fn condition_from_ssh_wire_ignores_an_unused_id() {
let ssh = |e, id| condition_from_wire(e, id, UnusedExpectedId::Ignore);
assert_eq!(ssh(1, ID_A), Ok(RefWriteCondition::Any));
assert_eq!(ssh(2, &[1, 2, 3]), Ok(RefWriteCondition::Missing));
assert_eq!(ssh(0, &[]), Err(RefWireError::Unspecified));
assert_eq!(
ssh(3, &[0; 31]).unwrap_err().ssh_message(),
"MATCH expectation requires a 32-byte expected_id"
);
}
#[test]
fn condition_from_wire_match_needs_32_bytes() {
assert_eq!(
rejected(3, &[0; 31]).connect_message(),
"expected a 32-byte digest, got 31 bytes"
);
assert_eq!(
rejected(3, &[]).connect_message(),
"expected a 32-byte digest, got 0 bytes"
);
}
#[test]
fn condition_from_wire_ok() {
let connect = |e, id| condition_from_wire(e, id, UnusedExpectedId::Reject);
assert_eq!(connect(1, &[]), Ok(RefWriteCondition::Any));
assert_eq!(connect(2, &[]), Ok(RefWriteCondition::Missing));
assert_eq!(connect(3, ID_B), Ok(RefWriteCondition::Match([0xbb; 32])));
}
#[test]
fn strip_listed_prefix_matches_at_component_boundaries() {
let strip = strip_listed_prefix;
for p in ["refs/heads", "refs/heads/", "refs/heads//"] {
assert_eq!(strip("refs/heads/main", p), Some("main"), "{p}");
assert_eq!(strip("refs/heads/feat/x", p), Some("feat/x"), "{p}");
assert_eq!(strip("refs/tags/v1", p), None, "{p}");
}
assert_eq!(strip("refs/heads/main", ""), Some("refs/heads/main"));
assert_eq!(strip("refs/heads/main", "refs//"), Some("heads/main"));
assert_eq!(strip("refs/heads/main", "refs"), Some("heads/main"));
assert_eq!(strip("refs/heads/main", "refs/heads/ma"), None);
assert_eq!(strip("refs/heads/featx", "refs/heads/feat"), None);
assert_eq!(strip("refs/heads/feat/x", "refs/heads/feat"), Some("x"));
assert_eq!(strip("refs/heads/main", "refs/heads/main"), None);
assert_eq!(list_scan_prefix(""), "");
assert_eq!(list_scan_prefix("refs//"), "refs/");
assert_eq!(list_scan_prefix("refs/heads/main"), "refs/heads/main/");
}
#[test]
fn ref_name_validation_is_mkit_core() {
assert!(validate_ref_name("refs/heads/main"));
assert!(!validate_ref_name("refs/heads/../main"));
let longest = format!("refs/heads/{}", "a".repeat(MAX_REF_NAME_BYTES - 11));
assert!(validate_ref_name(&longest));
assert!(!validate_ref_name(&format!("{longest}a")));
assert!(validate_ref_prefix(""));
assert!(validate_ref_prefix("refs/heads/"));
assert!(!validate_ref_prefix("/"));
}
}