Skip to main content

rill_runtime/
package.rs

1use std::{
2    collections::BTreeMap,
3    io::{Cursor, Read, Seek, Write},
4};
5
6use ed25519_dalek::{Signature, Signer, SigningKey, Verifier, VerifyingKey};
7use rill_runtime_protocol::{ModelPackManifest, ReleaseIndexPayload, SignedReleaseIndex};
8use serde::{Deserialize, Serialize};
9use serde_json::Value;
10use sha2::{Digest, Sha256};
11use thiserror::Error;
12use zip::{ZipArchive, ZipWriter, write::SimpleFileOptions};
13
14const MAX_FILES: usize = 8;
15const MAX_FILE_BYTES: u64 = 256 * 1024;
16const MAX_TOTAL_BYTES: u64 = 1024 * 1024;
17
18const MANIFEST_PATH: &str = "manifest.json";
19const MODEL_PATH: &str = "model.json";
20const CHECKSUMS_PATH: &str = "checksums.json";
21const SIGNATURE_PATH: &str = "META-INF/signature.ed25519";
22
23#[derive(Debug, Default, Clone)]
24pub struct TrustStore(pub BTreeMap<String, VerifyingKey>);
25
26#[derive(Debug, Clone)]
27pub struct LoadedModelPack {
28    pub manifest: ModelPackManifest,
29    pub model: serde_json::Value,
30}
31
32#[derive(Debug, Clone, Serialize)]
33#[serde(rename_all = "camelCase")]
34pub struct ModelPackInspection {
35    pub id: String,
36    pub version: String,
37    pub publisher_key_id: String,
38    pub runtime_api_version: u32,
39    pub capabilities: Vec<String>,
40    pub signature_verified: bool,
41}
42
43#[derive(Debug, Serialize, Deserialize)]
44#[serde(rename_all = "camelCase", deny_unknown_fields)]
45struct Checksums {
46    schema_version: u32,
47    files: BTreeMap<String, String>,
48}
49
50pub fn canonical_json(bytes: &[u8]) -> Result<Vec<u8>, ModelPackError> {
51    fn sort(value: Value) -> Value {
52        match value {
53            Value::Object(map) => Value::Object(
54                map.into_iter()
55                    .map(|(key, value)| (key, sort(value)))
56                    .collect(),
57            ),
58            Value::Array(items) => Value::Array(items.into_iter().map(sort).collect()),
59            other => other,
60        }
61    }
62    let value: Value = serde_json::from_slice(bytes)?;
63    Ok(serde_json::to_vec(&sort(value))?)
64}
65
66pub fn sign_release_index(
67    payload: ReleaseIndexPayload,
68    signing_key: &SigningKey,
69) -> Result<SignedReleaseIndex, ReleaseIndexError> {
70    validate_release_payload(&payload)?;
71    let serialized = serde_json::to_vec(&payload)?;
72    let canonical = canonical_json(&serialized).map_err(ReleaseIndexError::Canonical)?;
73    let signature = hex::encode(signing_key.sign(&canonical).to_bytes());
74    Ok(SignedReleaseIndex { payload, signature })
75}
76
77pub fn verify_release_index(
78    index: &SignedReleaseIndex,
79    trust: &TrustStore,
80) -> Result<(), ReleaseIndexError> {
81    validate_release_payload(&index.payload)?;
82    let signature_bytes =
83        hex::decode(&index.signature).map_err(|_| ReleaseIndexError::Signature)?;
84    let signature =
85        Signature::from_slice(&signature_bytes).map_err(|_| ReleaseIndexError::Signature)?;
86    let key = trust
87        .0
88        .get(&index.payload.publisher_key_id)
89        .ok_or(ReleaseIndexError::UnknownKey)?;
90    let serialized = serde_json::to_vec(&index.payload)?;
91    let canonical = canonical_json(&serialized).map_err(ReleaseIndexError::Canonical)?;
92    key.verify(&canonical, &signature)
93        .map_err(|_| ReleaseIndexError::Signature)
94}
95
96fn validate_release_payload(payload: &ReleaseIndexPayload) -> Result<(), ReleaseIndexError> {
97    payload
98        .validate_shape()
99        .map_err(|message| ReleaseIndexError::Manifest(message.into()))?;
100    let mut identities = std::collections::BTreeSet::new();
101    for artifact in &payload.artifacts {
102        semver::Version::parse(&artifact.version).map_err(|error| {
103            ReleaseIndexError::Manifest(format!("invalid artifact version: {error}"))
104        })?;
105        let identity = (
106            artifact.kind.clone(),
107            artifact.id.clone(),
108            artifact.target_os.clone(),
109            artifact.target_arch.clone(),
110        );
111        if !identities.insert(identity) {
112            return Err(ReleaseIndexError::Manifest(
113                "duplicate release artifact identity".into(),
114            ));
115        }
116    }
117    Ok(())
118}
119
120pub fn load_model_pack<R: Read + Seek>(
121    reader: R,
122    trust: &TrustStore,
123) -> Result<(LoadedModelPack, ModelPackInspection), ModelPackError> {
124    let files = read_archive(reader)?;
125    let manifest_bytes = files
126        .get(MANIFEST_PATH)
127        .ok_or(ModelPackError::Missing(MANIFEST_PATH))?;
128    let manifest: ModelPackManifest = serde_json::from_slice(manifest_bytes)?;
129    manifest
130        .validate_shape()
131        .map_err(|message| ModelPackError::Manifest(message.into()))?;
132    semver::Version::parse(&manifest.version)
133        .map_err(|error| ModelPackError::Manifest(format!("invalid pack version: {error}")))?;
134    let minimum = semver::Version::parse(&manifest.min_runtime_version)
135        .map_err(|error| ModelPackError::Manifest(format!("invalid minimum runtime: {error}")))?;
136    let runtime = semver::Version::parse(env!("CARGO_PKG_VERSION"))
137        .map_err(|error| ModelPackError::Manifest(format!("invalid runtime version: {error}")))?;
138    if runtime < minimum {
139        return Err(ModelPackError::RuntimeTooOld {
140            minimum: minimum.to_string(),
141            actual: runtime.to_string(),
142        });
143    }
144    verify_checksums_and_signature(&files, &manifest, trust)?;
145
146    let model: serde_json::Value = serde_json::from_slice(
147        files
148            .get(MODEL_PATH)
149            .ok_or(ModelPackError::Missing(MODEL_PATH))?,
150    )?;
151    let inspection = ModelPackInspection {
152        id: manifest.id.clone(),
153        version: manifest.version.clone(),
154        publisher_key_id: manifest.publisher_key_id.clone(),
155        runtime_api_version: manifest.runtime_api_version,
156        capabilities: manifest.capabilities.clone(),
157        signature_verified: true,
158    };
159    Ok((LoadedModelPack { manifest, model }, inspection))
160}
161
162fn read_archive<R: Read + Seek>(reader: R) -> Result<BTreeMap<String, Vec<u8>>, ModelPackError> {
163    let mut archive = ZipArchive::new(reader)?;
164    if archive.len() > MAX_FILES {
165        return Err(ModelPackError::Limit("file count"));
166    }
167    let mut total = 0u64;
168    let mut files = BTreeMap::new();
169    for index in 0..archive.len() {
170        let mut entry = archive.by_index(index)?;
171        if entry.is_dir() {
172            continue;
173        }
174        let name = entry.name().to_string();
175        validate_path(&name)?;
176        if !matches!(
177            name.as_str(),
178            MANIFEST_PATH | MODEL_PATH | CHECKSUMS_PATH | SIGNATURE_PATH
179        ) {
180            return Err(ModelPackError::Forbidden(name));
181        }
182        if entry.size() > MAX_FILE_BYTES {
183            return Err(ModelPackError::Limit("file size"));
184        }
185        total = total
186            .checked_add(entry.size())
187            .ok_or(ModelPackError::Limit("total size"))?;
188        if total > MAX_TOTAL_BYTES {
189            return Err(ModelPackError::Limit("total size"));
190        }
191        let mut bytes = Vec::with_capacity(entry.size() as usize);
192        entry.read_to_end(&mut bytes)?;
193        if files.insert(name.clone(), bytes).is_some() {
194            return Err(ModelPackError::Duplicate(name));
195        }
196    }
197    Ok(files)
198}
199
200fn verify_checksums_and_signature(
201    files: &BTreeMap<String, Vec<u8>>,
202    manifest: &ModelPackManifest,
203    trust: &TrustStore,
204) -> Result<(), ModelPackError> {
205    let checksum_bytes = files
206        .get(CHECKSUMS_PATH)
207        .ok_or(ModelPackError::Missing(CHECKSUMS_PATH))?;
208    let checksums: Checksums = serde_json::from_slice(checksum_bytes)?;
209    if checksums.schema_version != 1 {
210        return Err(ModelPackError::Manifest(
211            "unsupported checksum schema".into(),
212        ));
213    }
214    let expected_names = vec![MANIFEST_PATH.to_string(), MODEL_PATH.to_string()];
215    if checksums.files.keys().cloned().collect::<Vec<_>>() != expected_names {
216        return Err(ModelPackError::ChecksumCoverage);
217    }
218    for (name, expected) in &checksums.files {
219        let bytes = files
220            .get(name)
221            .ok_or_else(|| ModelPackError::MissingOwned(name.clone()))?;
222        let actual = hex::encode(Sha256::digest(bytes));
223        if &actual != expected {
224            return Err(ModelPackError::Digest(name.clone()));
225        }
226    }
227    let raw_signature = files
228        .get(SIGNATURE_PATH)
229        .ok_or(ModelPackError::Missing(SIGNATURE_PATH))?;
230    let signature = Signature::from_slice(raw_signature).map_err(|_| ModelPackError::Signature)?;
231    let key = trust
232        .0
233        .get(&manifest.publisher_key_id)
234        .ok_or(ModelPackError::UnknownKey)?;
235    let mut message = canonical_json(
236        files
237            .get(MANIFEST_PATH)
238            .ok_or(ModelPackError::Missing(MANIFEST_PATH))?,
239    )?;
240    message.push(b'\n');
241    message.extend(canonical_json(checksum_bytes)?);
242    key.verify(&message, &signature)
243        .map_err(|_| ModelPackError::Signature)
244}
245
246pub fn build_signed_model_pack(
247    manifest: &ModelPackManifest,
248    model: &serde_json::Value,
249    signing_key: &SigningKey,
250) -> Result<Vec<u8>, ModelPackError> {
251    manifest
252        .validate_shape()
253        .map_err(|message| ModelPackError::Manifest(message.into()))?;
254    let manifest_bytes = serde_json::to_vec_pretty(manifest)?;
255    let model_bytes = serde_json::to_vec_pretty(model)?;
256    let checksums = Checksums {
257        schema_version: 1,
258        files: BTreeMap::from([
259            (
260                MANIFEST_PATH.into(),
261                hex::encode(Sha256::digest(&manifest_bytes)),
262            ),
263            (MODEL_PATH.into(), hex::encode(Sha256::digest(&model_bytes))),
264        ]),
265    };
266    let checksum_bytes = serde_json::to_vec_pretty(&checksums)?;
267    let mut message = canonical_json(&manifest_bytes)?;
268    message.push(b'\n');
269    message.extend(canonical_json(&checksum_bytes)?);
270    let signature = signing_key.sign(&message).to_bytes();
271
272    let mut output = Cursor::new(Vec::new());
273    {
274        let mut archive = ZipWriter::new(&mut output);
275        let options = SimpleFileOptions::default()
276            .compression_method(zip::CompressionMethod::Deflated)
277            .unix_permissions(0o644);
278        for (name, bytes) in [
279            (MANIFEST_PATH, manifest_bytes.as_slice()),
280            (MODEL_PATH, model_bytes.as_slice()),
281            (CHECKSUMS_PATH, checksum_bytes.as_slice()),
282            (SIGNATURE_PATH, signature.as_slice()),
283        ] {
284            archive.start_file(name, options)?;
285            archive.write_all(bytes)?;
286        }
287        archive.finish()?;
288    }
289    Ok(output.into_inner())
290}
291
292fn validate_path(name: &str) -> Result<(), ModelPackError> {
293    if name.starts_with('/')
294        || name.contains('\\')
295        || name
296            .split('/')
297            .any(|part| part.is_empty() || part == "." || part == "..")
298    {
299        return Err(ModelPackError::UnsafePath(name.into()));
300    }
301    Ok(())
302}
303
304#[derive(Debug, Error)]
305pub enum ModelPackError {
306    #[error("zip error: {0}")]
307    Zip(#[from] zip::result::ZipError),
308    #[error("I/O error: {0}")]
309    Io(#[from] std::io::Error),
310    #[error("JSON error: {0}")]
311    Json(#[from] serde_json::Error),
312    #[error("unsafe package path {0}")]
313    UnsafePath(String),
314    #[error("forbidden package file {0}")]
315    Forbidden(String),
316    #[error("duplicate package file {0}")]
317    Duplicate(String),
318    #[error("model package exceeded {0} limit")]
319    Limit(&'static str),
320    #[error("missing package file {0}")]
321    Missing(&'static str),
322    #[error("missing package file {0}")]
323    MissingOwned(String),
324    #[error("invalid model manifest: {0}")]
325    Manifest(String),
326    #[error("checksum coverage does not exactly match the model payload")]
327    ChecksumCoverage,
328    #[error("checksum mismatch for {0}")]
329    Digest(String),
330    #[error("unknown publisher key")]
331    UnknownKey,
332    #[error("signature verification failed")]
333    Signature,
334    #[error("runtime {actual} is older than model requirement {minimum}")]
335    RuntimeTooOld { minimum: String, actual: String },
336}
337
338#[derive(Debug, Error)]
339pub enum ReleaseIndexError {
340    #[error("JSON error: {0}")]
341    Json(#[from] serde_json::Error),
342    #[error("invalid release index: {0}")]
343    Manifest(String),
344    #[error("unknown release-index publisher key")]
345    UnknownKey,
346    #[error("release-index signature verification failed")]
347    Signature,
348    #[error("canonical JSON error: {0}")]
349    Canonical(ModelPackError),
350}
351
352#[cfg(test)]
353mod tests {
354    use super::*;
355    use rill_runtime_protocol::{MODEL_PACK_FORMAT_VERSION, RUNTIME_API_VERSION};
356
357    fn manifest(key_id: &str) -> ModelPackManifest {
358        ModelPackManifest {
359            format_version: MODEL_PACK_FORMAT_VERSION,
360            id: "rillml.example.default".into(),
361            version: "0.5.0".into(),
362            runtime_api_version: RUNTIME_API_VERSION,
363            min_runtime_version: "0.5.0".into(),
364            publisher_key_id: key_id.into(),
365            capabilities: vec!["rillml.example".into()],
366        }
367    }
368
369    #[test]
370    fn signed_pack_roundtrip_and_tamper_rejection() {
371        let signing = SigningKey::from_bytes(&[7; 32]);
372        let key_id = "test-key";
373        let bytes = build_signed_model_pack(
374            &manifest(key_id),
375            &serde_json::json!({"description": "test"}),
376            &signing,
377        )
378        .unwrap();
379        let trust = TrustStore(BTreeMap::from([(key_id.into(), signing.verifying_key())]));
380        let (loaded, inspection) = load_model_pack(Cursor::new(&bytes), &trust).unwrap();
381        assert_eq!(loaded.manifest.id, "rillml.example.default");
382        assert!(inspection.signature_verified);
383
384        let wrong = SigningKey::from_bytes(&[8; 32]);
385        let wrong_trust = TrustStore(BTreeMap::from([(key_id.into(), wrong.verifying_key())]));
386        assert!(matches!(
387            load_model_pack(Cursor::new(bytes), &wrong_trust),
388            Err(ModelPackError::Signature)
389        ));
390    }
391
392    #[test]
393    fn release_index_signature_covers_artifact_hashes() {
394        use rill_runtime_protocol::{
395            RELEASE_INDEX_SCHEMA_VERSION, RUNTIME_ARTIFACT_ID, ReleaseArtifact,
396            ReleaseArtifactKind, ReleaseIndexPayload,
397        };
398
399        let signing = SigningKey::from_bytes(&[6; 32]);
400        let payload = ReleaseIndexPayload {
401            schema_version: RELEASE_INDEX_SCHEMA_VERSION,
402            channel: "stable".into(),
403            generated_at: "2026-07-13T00:00:00Z".into(),
404            publisher_key_id: "release-test".into(),
405            artifacts: vec![ReleaseArtifact {
406                kind: ReleaseArtifactKind::Runtime,
407                id: RUNTIME_ARTIFACT_ID.into(),
408                version: "0.5.0".into(),
409                runtime_api_version: RUNTIME_API_VERSION,
410                target_os: Some("macos".into()),
411                target_arch: Some("aarch64".into()),
412                url: "https://example.invalid/rill-runtime".into(),
413                sha256: "ab".repeat(32),
414                size: 1024,
415            }],
416        };
417        let mut index = sign_release_index(payload, &signing).unwrap();
418        let trust = TrustStore(BTreeMap::from([(
419            "release-test".into(),
420            signing.verifying_key(),
421        )]));
422        verify_release_index(&index, &trust).unwrap();
423        index.payload.artifacts[0].sha256 = "cd".repeat(32);
424        assert!(matches!(
425            verify_release_index(&index, &trust),
426            Err(ReleaseIndexError::Signature)
427        ));
428    }
429}