1use std::collections::BTreeSet;
4
5use mkit_core::hash::Hash;
6use serde::{Deserialize, Serialize};
7
8use super::outbox::guard;
9use super::{Batch, BatchOutcome, NamespaceStore, Partition, StoreError, Value, codec, keys};
10use crate::pipeline::ShardMap;
11use crate::repo::RepoId;
12
13pub const MAX_FLAG_IDS: usize = 48;
15pub const MAX_FLAG_SOURCES: usize = 1024;
17const MAX_ATTEMPTS: usize = 16;
18
19#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
21#[serde(deny_unknown_fields)]
22pub struct FlagSource {
23 pub inspector: String,
25 pub inspection_id: String,
27 pub ref_name: String,
29 pub sequence: u64,
31}
32
33#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
35#[serde(rename_all = "lowercase")]
36pub enum FlagState {
37 Flagged,
39 Released,
41}
42
43#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
45#[serde(deny_unknown_fields)]
46pub struct FlagV1 {
47 pub id: Hash,
49 pub reason: String,
51 pub source: FlagSource,
53 pub state: FlagState,
55 pub seen_sources: Vec<Hash>,
57}
58
59impl FlagV1 {
60 pub fn new(request: FlagInstall) -> Result<Self, StoreError> {
63 if !valid(&request.reason, &request.source) {
64 return Err(StoreError::Invalid("invalid inspection flag".into()));
65 }
66 let seen_sources = vec![source_id(&request.source)];
67 Ok(Self {
68 id: request.id,
69 reason: request.reason,
70 source: request.source,
71 state: FlagState::Flagged,
72 seen_sources,
73 })
74 }
75}
76
77fn source_id(source: &FlagSource) -> Hash {
78 mkit_core::hash::hash(&serde_json::to_vec(source).expect("FlagSource is JSON serializable"))
79}
80
81fn valid_history(record: &FlagV1) -> bool {
82 record.seen_sources.len() <= MAX_FLAG_SOURCES
83 && record.seen_sources.windows(2).all(|pair| pair[0] < pair[1])
84 && record
85 .seen_sources
86 .binary_search(&source_id(&record.source))
87 .is_ok()
88}
89
90fn valid(reason: &str, source: &FlagSource) -> bool {
91 let nonempty = |s: &str, max: usize| !s.is_empty() && s.len() <= max && !s.contains('\0');
92 nonempty(reason, 4096)
93 && nonempty(&source.inspector, 1024)
94 && nonempty(&source.inspection_id, 1024)
95 && source.ref_name.len() <= 1024
96 && crate::refs::is_served_ref_name(&source.ref_name)
97 && source.sequence > 0
98}
99
100pub fn encode_flag(record: &FlagV1) -> Result<Value, StoreError> {
103 if !valid(&record.reason, &record.source) || !valid_history(record) {
104 return Err(StoreError::Invalid("invalid inspection flag".into()));
105 }
106 let mut bytes = vec![codec::CODEC_V1];
107 serde_json::to_writer(&mut bytes, record)
108 .map_err(|_| StoreError::Invalid("invalid inspection flag".into()))?;
109 Ok(Value::new(bytes))
110}
111
112pub fn decode_flag(value: &Value) -> Result<FlagV1, StoreError> {
115 let corrupt = || StoreError::Corrupt("invalid inspection flag".into());
116 let Some((&codec::CODEC_V1, body)) = value.as_bytes().split_first() else {
117 return Err(corrupt());
118 };
119 let record: FlagV1 = serde_json::from_slice(body).map_err(|_| corrupt())?;
120 if !valid(&record.reason, &record.source) || !valid_history(&record) {
121 return Err(corrupt());
122 }
123 Ok(record)
124}
125
126#[derive(Debug, Clone)]
128pub struct FlagInstall {
129 pub id: Hash,
131 pub reason: String,
133 pub source: FlagSource,
135}
136
137#[derive(Debug, Clone, PartialEq, Eq)]
139pub struct FlagLookup {
140 pub flagged: Vec<FlagV1>,
142 pub version: u64,
144}
145
146#[derive(Debug)]
148pub struct InspectionFlags<'a, S: ?Sized> {
149 store: &'a S,
150 repo: &'a RepoId,
151 partition: Partition,
152}
153
154impl<'a, S: NamespaceStore + ?Sized> InspectionFlags<'a, S> {
155 #[must_use]
157 pub fn new(store: &'a S, shards: &dyn ShardMap, repo: &'a RepoId) -> Self {
158 Self {
159 store,
160 repo,
161 partition: shards.object_index(repo, &[0; 32]),
162 }
163 }
164
165 async fn read_version(&self) -> Result<(Option<Value>, u64), StoreError> {
166 let value = self
167 .store
168 .get(&self.partition, &keys::inspection_version(&self.repo.name))
169 .await?;
170 let version = value
171 .as_ref()
172 .map(codec::decode_u64)
173 .transpose()?
174 .unwrap_or(0);
175 if value.is_some() && version == 0 {
176 return Err(StoreError::Corrupt(
177 "inspection registry version is zero".into(),
178 ));
179 }
180 Ok((value, version))
181 }
182
183 async fn records(&self, ids: &[Hash]) -> Result<Vec<Option<Value>>, StoreError> {
184 let keys: Vec<_> = ids
185 .iter()
186 .map(|id| keys::inspection_flag(&self.repo.name, id))
187 .collect();
188 self.store.get_many(&self.partition, &keys).await
189 }
190
191 fn record(id: &Hash, value: Option<&Value>) -> Result<Option<FlagV1>, StoreError> {
192 let record = value.map(decode_flag).transpose()?;
193 if record.as_ref().is_some_and(|r| r.id != *id) {
194 return Err(StoreError::Corrupt(
195 "inspection flag key binding mismatch".into(),
196 ));
197 }
198 Ok(record)
199 }
200
201 async fn commit(&self, batch: Batch) -> Result<bool, StoreError> {
202 batch.validate(&self.store.capabilities())?;
203 match self.store.apply(&self.partition, batch).await? {
204 BatchOutcome::Committed => Ok(true),
205 BatchOutcome::PreconditionFailed { .. } => Ok(false),
206 BatchOutcome::DeadlinePassed { .. } => {
207 Err(StoreError::Corrupt("unexpected inspection deadline".into()))
208 }
209 }
210 }
211
212 pub async fn install_flags(&self, flags: &[FlagInstall]) -> Result<u64, StoreError> {
216 if flags.len() > MAX_FLAG_IDS || flags.iter().any(|f| !valid(&f.reason, &f.source)) {
217 return Err(StoreError::Invalid(
218 "invalid or oversized inspection flags".into(),
219 ));
220 }
221 let ids: Vec<_> = flags.iter().map(|f| f.id).collect();
222 bounded(&ids)?;
223 let sources: Vec<_> = flags.iter().map(|f| source_id(&f.source)).collect();
224 self.mutate(&ids, |index, prior| {
225 let request = &flags[index];
226 let Some(prior) = prior else {
227 return encode_flag(&FlagV1::new(request.clone())?).map(Some);
228 };
229 let Err(position) = prior.seen_sources.binary_search(&sources[index]) else {
230 return Ok(None);
231 };
232 if prior.seen_sources.len() == MAX_FLAG_SOURCES {
233 return Err(StoreError::Invalid(
234 "inspection flag source history exhausted".into(),
235 ));
236 }
237 let mut next = prior.clone();
238 next.seen_sources.insert(position, sources[index]);
239 if next.state == FlagState::Released {
240 next.reason.clone_from(&request.reason);
241 next.source.clone_from(&request.source);
242 next.state = FlagState::Flagged;
243 }
244 encode_flag(&next).map(Some)
245 })
246 .await
247 }
248
249 pub async fn release_flag(&self, id: &Hash) -> Result<u64, StoreError> {
252 self.mutate(&[*id], |_, prior| {
253 let Some(record) = prior.filter(|r| r.state == FlagState::Flagged) else {
254 return Ok(None);
255 };
256 let mut released = record.clone();
257 released.state = FlagState::Released;
258 encode_flag(&released).map(Some)
259 })
260 .await
261 }
262
263 async fn mutate(
264 &self,
265 ids: &[Hash],
266 replacement: impl Fn(usize, Option<&FlagV1>) -> Result<Option<Value>, StoreError>
267 + crate::rt::MaybeSync,
268 ) -> Result<u64, StoreError> {
269 for _ in 0..MAX_ATTEMPTS {
270 let (version_raw, mut version) = self.read_version().await?;
271 let values = self.records(ids).await?;
272 let version_key = keys::inspection_version(&self.repo.name);
273 if version_raw.is_none() && values.iter().any(Option::is_some) {
274 if self
275 .commit(Batch::new().require(guard(version_key, None)))
276 .await?
277 {
278 return Err(StoreError::Corrupt(
279 "inspection flags without registry version".into(),
280 ));
281 }
282 continue;
283 }
284 let mut batch = Batch::new().require(guard(version_key.clone(), version_raw.as_ref()));
285 for (index, (id, value)) in ids.iter().zip(&values).enumerate() {
286 let record = Self::record(id, value.as_ref())?;
287 let key = keys::inspection_flag(&self.repo.name, id);
288 batch.preconditions.push(guard(key.clone(), value.as_ref()));
289 if let Some(next) = replacement(index, record.as_ref())? {
290 version = version.checked_add(1).ok_or_else(|| {
291 StoreError::Corrupt("inspection registry version overflow".into())
292 })?;
293 batch = batch.put(key, next);
294 }
295 }
296 if !batch.writes.is_empty() {
297 batch = batch.put(version_key, codec::encode_u64(version));
298 }
299 if self.commit(batch).await? {
300 return Ok(version);
301 }
302 }
303 Err(contended())
304 }
305
306 pub async fn lookup(&self, ids: &[Hash]) -> Result<FlagLookup, StoreError> {
309 bounded(ids)?;
310 for _ in 0..MAX_ATTEMPTS {
311 let (raw, version) = self.read_version().await?;
312 let values = self.records(ids).await?;
313 let mut flagged = Vec::new();
314 for (id, value) in ids.iter().zip(&values) {
315 if let Some(record) =
316 Self::record(id, value.as_ref())?.filter(|r| r.state == FlagState::Flagged)
317 {
318 flagged.push(record);
319 }
320 }
321 let batch = Batch::new().require(guard(
322 keys::inspection_version(&self.repo.name),
323 raw.as_ref(),
324 ));
325 if self.commit(batch).await? {
326 if raw.is_none() && values.iter().any(Option::is_some) {
327 return Err(StoreError::Corrupt(
328 "inspection flags without registry version".into(),
329 ));
330 }
331 return Ok(FlagLookup { flagged, version });
332 }
333 }
334 Err(contended())
335 }
336
337 pub async fn version(&self) -> Result<u64, StoreError> {
340 self.read_version().await.map(|(_, version)| version)
341 }
342}
343
344fn bounded(ids: &[Hash]) -> Result<(), StoreError> {
345 if ids.len() > MAX_FLAG_IDS || ids.iter().collect::<BTreeSet<_>>().len() != ids.len() {
346 return Err(StoreError::Invalid(
347 "inspection ids must be distinct and bounded".into(),
348 ));
349 }
350 Ok(())
351}
352
353fn contended() -> StoreError {
354 StoreError::unavailable(std::io::Error::other("inspection registry CAS retry limit"))
355}
356
357#[cfg(test)]
358#[path = "inspection_flags_tests.rs"]
359mod tests;