1use bytes::{Bytes, BytesMut};
4use notedthat_core::{ConditionalHeaders, KbSlug, ObjectPath, Storage, StorageError};
5
6use crate::sinks::WriteSinks;
7use crate::{ReplaceOutcome, WriteError};
8
9pub struct ReplaceRequest<'a> {
11 pub kb: &'a KbSlug,
13 pub path: &'a ObjectPath,
15 pub old_string: &'a str,
17 pub new_string: &'a str,
19 pub replace_all: bool,
21 pub caller_conditionals: ConditionalHeaders,
23 pub max_patchable_size: u64,
25 pub caller_content_type: Option<&'a str>,
27}
28
29pub 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;