1use super::{AprV2Reader, AprV2Writer, V2FormatError};
23
24#[derive(Debug, Clone, Default)]
40pub struct ProvenancePatch {
41 pub license: Option<String>,
43 pub data_source: Option<String>,
45 pub data_license: Option<String>,
47 pub hf_architecture: Option<String>,
50 pub hf_model_type: Option<String>,
53 pub architecture: Option<String>,
59 pub tokenizer_vocab: Option<Vec<String>>,
73 pub tokenizer_merges: Option<Vec<String>>,
76 pub tokenizer_model_type: Option<String>,
79}
80
81impl ProvenancePatch {
82 #[must_use]
85 pub fn has_any(&self) -> bool {
86 self.license.is_some()
87 || self.data_source.is_some()
88 || self.data_license.is_some()
89 || self.hf_architecture.is_some()
90 || self.hf_model_type.is_some()
91 || self.architecture.is_some()
92 || self.tokenizer_vocab.is_some()
93 || self.tokenizer_merges.is_some()
94 || self.tokenizer_model_type.is_some()
95 }
96}
97
98pub fn stamp_provenance_bytes(
121 input: &[u8],
122 patch: &ProvenancePatch,
123) -> Result<Vec<u8>, V2FormatError> {
124 if !patch.has_any() {
125 return Err(V2FormatError::InvalidHeader(
126 "stamp_provenance_bytes: patch has no fields set — \
127 refusing to rewrite without changes"
128 .to_string(),
129 ));
130 }
131
132 let reader = AprV2Reader::from_bytes(input)?;
133
134 let original_flags = reader.header().flags;
135 let mut new_metadata = reader.metadata().clone();
136
137 if let Some(ref lic) = patch.license {
138 new_metadata.license = Some(lic.clone());
139 }
140 if let Some(ref ds) = patch.data_source {
141 new_metadata.data_source = Some(ds.clone());
142 }
143 if let Some(ref dl) = patch.data_license {
144 new_metadata.data_license = Some(dl.clone());
145 }
146 if let Some(ref ha) = patch.hf_architecture {
148 new_metadata.hf_architecture = Some(ha.clone());
149 }
150 if let Some(ref hmt) = patch.hf_model_type {
151 new_metadata.hf_model_type = Some(hmt.clone());
152 }
153 if let Some(ref arch) = patch.architecture {
154 new_metadata.architecture = Some(arch.clone());
155 }
156 let mut set_has_vocab = false;
162 if let Some(ref vocab) = patch.tokenizer_vocab {
163 if !vocab.is_empty() {
164 let vocab_array: Vec<serde_json::Value> = vocab
165 .iter()
166 .map(|s| serde_json::Value::String(s.clone()))
167 .collect();
168 new_metadata.custom.insert(
169 "tokenizer.vocabulary".to_string(),
170 serde_json::Value::Array(vocab_array),
171 );
172 new_metadata.custom.insert(
173 "tokenizer.vocab_size".to_string(),
174 serde_json::Value::Number(serde_json::Number::from(vocab.len())),
175 );
176 set_has_vocab = true;
177 }
178 }
179 if let Some(ref merges) = patch.tokenizer_merges {
180 if !merges.is_empty() {
181 let merges_array: Vec<serde_json::Value> = merges
182 .iter()
183 .map(|s| serde_json::Value::String(s.clone()))
184 .collect();
185 new_metadata.custom.insert(
186 "tokenizer.merges".to_string(),
187 serde_json::Value::Array(merges_array),
188 );
189 }
190 }
191 if let Some(ref mt) = patch.tokenizer_model_type {
192 new_metadata.custom.insert(
193 "tokenizer.model_type".to_string(),
194 serde_json::Value::String(mt.clone()),
195 );
196 }
197
198 let effective_flags = if set_has_vocab {
203 original_flags.with(super::AprV2Flags::HAS_VOCAB)
204 } else {
205 original_flags
206 };
207
208 let mut writer = AprV2Writer::new(new_metadata);
209 writer.set_header_flags(effective_flags);
210
211 for name in reader.tensor_names() {
214 let entry = reader
215 .get_tensor(name)
216 .ok_or_else(|| V2FormatError::InvalidHeader(format!("tensor {name} vanished")))?;
217 let data = reader
218 .get_tensor_data(name)
219 .ok_or_else(|| V2FormatError::InvalidHeader(format!("tensor {name} has no data")))?;
220 writer.add_tensor(
221 name.to_string(),
222 entry.dtype,
223 entry.shape.clone(),
224 data.to_vec(),
225 );
226 }
227
228 writer.write()
229}
230
231#[cfg(test)]
232mod tests {
233 use super::super::{AprV2Flags, AprV2Metadata, TensorDType};
234 use super::*;
235
236 fn minimal_apr_with_flags(flags: u16) -> Vec<u8> {
238 let metadata = AprV2Metadata::new("stamp-test");
239 let mut writer = AprV2Writer::new(metadata);
240 writer.set_header_flags(AprV2Flags::from_bits(flags));
241 writer.add_tensor(
242 "weight",
243 TensorDType::F32,
244 vec![2, 3],
245 vec![0u8; 24], );
247 writer.write().expect("write test apr")
248 }
249
250 #[test]
251 fn stamp_populates_all_three_fields_when_source_is_unpopulated() {
252 let input = minimal_apr_with_flags(0);
253 let patch = ProvenancePatch {
254 license: Some("Apache-2.0".into()),
255 data_source: Some("huggingface.co/Qwen/Qwen2.5-Coder-7B-Instruct".into()),
256 data_license: Some("Qwen-License-Agreement-v1".into()),
257 hf_architecture: None,
258 hf_model_type: None,
259 architecture: None,
260 tokenizer_vocab: None,
261 tokenizer_merges: None,
262 tokenizer_model_type: None,
263 };
264
265 let output = stamp_provenance_bytes(&input, &patch).expect("stamp must succeed");
266
267 let reader = AprV2Reader::from_bytes(&output).expect("stamped buffer must parse");
268 let md = reader.metadata();
269 assert_eq!(md.license.as_deref(), Some("Apache-2.0"));
270 assert_eq!(
271 md.data_source.as_deref(),
272 Some("huggingface.co/Qwen/Qwen2.5-Coder-7B-Instruct")
273 );
274 assert_eq!(
275 md.data_license.as_deref(),
276 Some("Qwen-License-Agreement-v1")
277 );
278 }
279
280 #[test]
281 fn stamp_preserves_tensor_data_byte_for_byte() {
282 let input = minimal_apr_with_flags(0);
283 let input_reader = AprV2Reader::from_bytes(&input).unwrap();
284 let original_bytes: Vec<u8> = input_reader
285 .get_tensor_data("weight")
286 .expect("input has weight")
287 .to_vec();
288
289 let patch = ProvenancePatch {
290 license: Some("MIT".into()),
291 ..Default::default()
292 };
293 let output = stamp_provenance_bytes(&input, &patch).unwrap();
294
295 let out_reader = AprV2Reader::from_bytes(&output).unwrap();
296 let round_tripped = out_reader
297 .get_tensor_data("weight")
298 .expect("output has weight");
299
300 assert_eq!(
301 original_bytes.as_slice(),
302 round_tripped,
303 "tensor bytes must survive stamp verbatim"
304 );
305 }
306
307 #[test]
308 fn stamp_preserves_header_flags() {
309 let flags = AprV2Flags::QUANTIZED | AprV2Flags::HAS_VOCAB;
311 let input = minimal_apr_with_flags(flags);
312
313 let in_reader = AprV2Reader::from_bytes(&input).unwrap();
314 assert!(in_reader.header().flags.contains(AprV2Flags::QUANTIZED));
315 assert!(in_reader.header().flags.contains(AprV2Flags::HAS_VOCAB));
316
317 let patch = ProvenancePatch {
318 license: Some("Apache-2.0".into()),
319 ..Default::default()
320 };
321 let output = stamp_provenance_bytes(&input, &patch).unwrap();
322
323 let out_reader = AprV2Reader::from_bytes(&output).unwrap();
324 assert!(
326 out_reader.header().flags.contains(AprV2Flags::QUANTIZED),
327 "QUANTIZED flag must survive stamp"
328 );
329 assert!(
330 out_reader.header().flags.contains(AprV2Flags::HAS_VOCAB),
331 "HAS_VOCAB flag must survive stamp"
332 );
333 assert!(
335 out_reader
336 .header()
337 .flags
338 .contains(AprV2Flags::LAYOUT_ROW_MAJOR),
339 "LAYOUT_ROW_MAJOR must always be set"
340 );
341 }
342
343 #[test]
344 fn stamp_rejects_empty_patch() {
345 let input = minimal_apr_with_flags(0);
346 let empty = ProvenancePatch::default();
347 let err = stamp_provenance_bytes(&input, &empty).unwrap_err();
348 let msg = format!("{err:?}");
349 assert!(
350 msg.contains("patch has no fields"),
351 "empty-patch error must be explicit: {msg}"
352 );
353 }
354
355 #[test]
356 fn stamp_allows_partial_patch_leaving_other_fields_unchanged() {
357 let mut md = AprV2Metadata::new("partial-test");
359 md.license = Some("Apache-2.0".into());
360 let mut writer = AprV2Writer::new(md);
361 writer.add_tensor("w", TensorDType::F32, vec![4], vec![0u8; 16]);
362 let input = writer.write().unwrap();
363
364 let patch = ProvenancePatch {
366 data_source: Some("teacher-only".into()),
367 ..Default::default()
368 };
369 let output = stamp_provenance_bytes(&input, &patch).unwrap();
370
371 let out_reader = AprV2Reader::from_bytes(&output).unwrap();
372 assert_eq!(
373 out_reader.metadata().license.as_deref(),
374 Some("Apache-2.0"),
375 "unchanged license must survive"
376 );
377 assert_eq!(
378 out_reader.metadata().data_source.as_deref(),
379 Some("teacher-only"),
380 "patched data_source must land"
381 );
382 assert!(
383 out_reader.metadata().data_license.is_none(),
384 "untouched data_license must remain None"
385 );
386 }
387
388 #[test]
389 fn stamp_is_idempotent_under_identical_patch() {
390 let input = minimal_apr_with_flags(0);
391 let patch = ProvenancePatch {
392 license: Some("Apache-2.0".into()),
393 data_source: Some("teacher-only".into()),
394 data_license: Some("Apache-2.0".into()),
395 hf_architecture: None,
396 hf_model_type: None,
397 architecture: None,
398 tokenizer_vocab: None,
399 tokenizer_merges: None,
400 tokenizer_model_type: None,
401 };
402
403 let first = stamp_provenance_bytes(&input, &patch).unwrap();
404 let second = stamp_provenance_bytes(&first, &patch).unwrap();
405 assert_eq!(
406 first, second,
407 "applying the same patch twice must be byte-identical (idempotent)"
408 );
409 }
410}