1use 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
10pub struct ReplaceRequest<'a> {
12 pub kb: &'a KbSlug,
14 pub path: &'a ObjectPath,
16 pub old_string: &'a str,
18 pub new_string: &'a str,
20 pub replace_all: bool,
22 pub caller_conditionals: ConditionalHeaders,
24 pub max_patchable_size: u64,
26 pub caller_content_type: Option<&'a str>,
28}
29
30pub 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;