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