Skip to main content

notedthat_write/
replace.rs

1//! Shared exact-substring replacement primitive for object writes.
2
3use bytes::{Bytes, BytesMut};
4use notedthat_core::{ConditionalHeaders, KbSlug, ObjectPath, Storage, StorageError};
5
6use crate::sinks::WriteSinks;
7use crate::{ReplaceOutcome, WriteError};
8
9/// Request data for one optimistic replace operation.
10pub struct ReplaceRequest<'a> {
11    /// Knowledge base containing the object.
12    pub kb: &'a KbSlug,
13    /// Object path to replace within.
14    pub path: &'a ObjectPath,
15    /// Exact UTF-8 substring to search for in the object body.
16    pub old_string: &'a str,
17    /// Replacement UTF-8 string to splice into the object body.
18    pub new_string: &'a str,
19    /// Whether to replace every non-overlapping match instead of exactly one match.
20    pub replace_all: bool,
21    /// Caller-supplied conditional headers.
22    pub caller_conditionals: ConditionalHeaders,
23    /// Maximum replaceable object size in bytes.
24    pub max_patchable_size: u64,
25    /// Caller-supplied content type, if any.
26    pub caller_content_type: Option<&'a str>,
27}
28
29/// Replace exact occurrences of a UTF-8 substring in an object body.
30///
31/// # Errors
32/// Returns [`crate::WriteError`] when storage access, preconditions, size limits, or
33/// match-count requirements fail.
34pub async fn replace(
35    storage: &dyn Storage,
36    sinks: &WriteSinks<'_>,
37    request: ReplaceRequest<'_>,
38) -> Result<ReplaceOutcome, WriteError> {
39    const MAX_ATTEMPTS: u32 = 3;
40    let ReplaceRequest {
41        kb,
42        path,
43        old_string,
44        new_string,
45        replace_all,
46        caller_conditionals,
47        max_patchable_size,
48        caller_content_type,
49    } = request;
50
51    if old_string.is_empty() {
52        return Err(WriteError::PatchInvalidRange {
53            message: "old_string must be non-empty (would match every byte position)".into(),
54        });
55    }
56    crate::patch::require_strong_if_match(&caller_conditionals)?;
57
58    let mut attempt = 0u32;
59    loop {
60        attempt += 1;
61
62        let meta = storage
63            .head_object(kb, path, ConditionalHeaders::default())
64            .await?;
65        ensure_caller_etag_matches(&caller_conditionals, meta.etag.as_deref())?;
66        ensure_within_patchable_size(meta.size, max_patchable_size)?;
67
68        let head_etag = meta
69            .etag
70            .clone()
71            .ok_or_else(|| WriteError::PatchInvalidRange {
72                message: "backend did not return ETag on HEAD".into(),
73            })?;
74
75        let get_conditionals = ConditionalHeaders {
76            if_match: Some(head_etag.clone()),
77            ..ConditionalHeaders::default()
78        };
79        let read = match storage.get_object(kb, path, None, get_conditionals).await {
80            Ok(read) => read,
81            Err(StorageError::PreconditionFailed) if attempt < MAX_ATTEMPTS => {
82                tracing::debug!(target: "notedthat::replace", kb = %kb, path = %path, attempt, stage = "get", "REPLACE_RETRY_PRECONDITION");
83                continue;
84            }
85            Err(error) => return Err(WriteError::Storage(error)),
86        };
87
88        let needle = old_string.as_bytes();
89        let haystack: &[u8] = read.bytes.as_ref();
90        let matches = find_non_overlapping_matches(haystack, needle);
91
92        if matches.is_empty() {
93            return Err(WriteError::ReplaceNoMatch);
94        }
95        if matches.len() >= 2 && !replace_all {
96            let count =
97                u64::try_from(matches.len()).map_err(|_| WriteError::PatchInvalidRange {
98                    message: "replace: match count exceeds u64".into(),
99                })?;
100            return Err(WriteError::ReplaceAmbiguous { count });
101        }
102
103        let match_bound = if replace_all { matches.len() } else { 1 };
104        let new_bytes = splice_replacement(&ReplacementSplice {
105            haystack,
106            needle,
107            replacement: new_string.as_bytes(),
108            matches: &matches,
109            match_bound,
110            max_patchable_size,
111        })?;
112
113        crate::manifest::check_manifest_bytes(kb, path, &new_bytes)?;
114        let put_conditionals = ConditionalHeaders {
115            if_match: Some(head_etag),
116            ..ConditionalHeaders::default()
117        };
118        let content_type = caller_content_type
119            .or(read.meta.content_type.as_deref())
120            .unwrap_or("application/octet-stream");
121        let new_size =
122            u64::try_from(new_bytes.len()).map_err(|_| WriteError::PatchInvalidRange {
123                message: "replace: body length exceeds u64".into(),
124            })?;
125        let put_outcome = match storage
126            .put_object(kb, path, new_bytes, Some(content_type), put_conditionals)
127            .await
128        {
129            Ok(outcome) => outcome,
130            Err(StorageError::PreconditionFailed) if attempt < MAX_ATTEMPTS => {
131                tracing::debug!(target: "notedthat::replace", kb = %kb, path = %path, attempt, stage = "put", "REPLACE_RETRY_PRECONDITION");
132                continue;
133            }
134            Err(error) => return Err(WriteError::Storage(error)),
135        };
136
137        crate::commit::after_write(sinks, kb, path, &put_outcome, new_size, content_type).await?;
138
139        let match_count =
140            u64::try_from(match_bound).map_err(|_| WriteError::PatchInvalidRange {
141                message: "replace: match count exceeds u64".into(),
142            })?;
143        return Ok(ReplaceOutcome {
144            put_outcome,
145            match_count,
146        });
147    }
148}
149
150fn ensure_caller_etag_matches(
151    caller_conditionals: &ConditionalHeaders,
152    current_etag: Option<&str>,
153) -> Result<(), WriteError> {
154    if caller_conditionals.if_match.as_deref() != current_etag {
155        return Err(WriteError::Storage(StorageError::PreconditionFailed));
156    }
157    Ok(())
158}
159
160fn ensure_within_patchable_size(size: u64, limit: u64) -> Result<(), WriteError> {
161    if size > limit {
162        return Err(WriteError::PatchTooLarge { size, limit });
163    }
164    Ok(())
165}
166
167fn find_non_overlapping_matches(haystack: &[u8], needle: &[u8]) -> Vec<usize> {
168    let mut matches = Vec::new();
169    let mut cursor = 0usize;
170    while cursor + needle.len() <= haystack.len() {
171        if &haystack[cursor..cursor + needle.len()] == needle {
172            matches.push(cursor);
173            cursor += needle.len();
174        } else {
175            cursor += 1;
176        }
177    }
178    matches
179}
180
181struct ReplacementSplice<'a> {
182    haystack: &'a [u8],
183    needle: &'a [u8],
184    replacement: &'a [u8],
185    matches: &'a [usize],
186    match_bound: usize,
187    max_patchable_size: u64,
188}
189
190fn splice_replacement(splice: &ReplacementSplice<'_>) -> Result<Bytes, WriteError> {
191    let replaced_bytes = splice.match_bound * splice.needle.len();
192    let added_bytes = splice.match_bound * splice.replacement.len();
193    let new_len = splice
194        .haystack
195        .len()
196        .checked_sub(replaced_bytes)
197        .and_then(|n| n.checked_add(added_bytes))
198        .ok_or_else(|| WriteError::PatchInvalidRange {
199            message: "replace: length arithmetic overflow".into(),
200        })?;
201    let new_len_u64 = u64::try_from(new_len).map_err(|_| WriteError::PatchInvalidRange {
202        message: "replace: new_len exceeds u64".into(),
203    })?;
204    if new_len_u64 > splice.max_patchable_size {
205        return Err(WriteError::PatchTooLarge {
206            size: new_len_u64,
207            limit: splice.max_patchable_size,
208        });
209    }
210
211    let mut result = BytesMut::with_capacity(new_len);
212    let mut prev_end = 0usize;
213    for &matched_at in splice.matches.iter().take(splice.match_bound) {
214        result.extend_from_slice(&splice.haystack[prev_end..matched_at]);
215        result.extend_from_slice(splice.replacement);
216        prev_end = matched_at + splice.needle.len();
217    }
218    result.extend_from_slice(&splice.haystack[prev_end..]);
219    Ok(result.freeze())
220}
221
222#[cfg(test)]
223mod tests;