Skip to main content

systemprompt_loader/bundle/source/oci/
pull.rs

1//! Manifest and blob reads against an OCI registry.
2//!
3//! Copyright (c) systemprompt.io — Business Source License 1.1.
4//! See <https://systemprompt.io> for licensing details.
5
6use std::path::Path;
7
8use serde::{Deserialize, Serialize};
9use sha2::{Digest, Sha256};
10use systemprompt_models::services::bundle::BUNDLE_MEDIA_TYPE;
11
12use super::RegistryClient;
13use crate::bundle::error::{BundleError, BundleResult};
14use crate::bundle::source::FetchedBundle;
15use crate::bundle::source::stream::stream_to_file;
16
17pub const OCI_MANIFEST_MEDIA_TYPE: &str = "application/vnd.oci.image.manifest.v1+json";
18pub const DOCKER_MANIFEST_MEDIA_TYPE: &str = "application/vnd.docker.distribution.manifest.v2+json";
19
20#[derive(Debug, Clone, Serialize, Deserialize)]
21pub struct OciDescriptor {
22    #[serde(rename = "mediaType")]
23    pub media_type: String,
24
25    pub digest: String,
26
27    pub size: u64,
28}
29
30#[derive(Debug, Clone, Serialize, Deserialize)]
31pub struct OciManifest {
32    #[serde(rename = "schemaVersion")]
33    pub schema_version: u32,
34
35    #[serde(rename = "mediaType", default, skip_serializing_if = "Option::is_none")]
36    pub media_type: Option<String>,
37
38    #[serde(
39        rename = "artifactType",
40        default,
41        skip_serializing_if = "Option::is_none"
42    )]
43    pub artifact_type: Option<String>,
44
45    pub config: OciDescriptor,
46
47    #[serde(default)]
48    pub layers: Vec<OciDescriptor>,
49}
50
51pub async fn get_manifest(registry: &RegistryClient) -> BundleResult<(String, OciManifest)> {
52    let url = registry.url(&format!("/manifests/{}", registry.manifest_ref()))?;
53    let accept = format!("{OCI_MANIFEST_MEDIA_TYPE}, {DOCKER_MANIFEST_MEDIA_TYPE}");
54    let response = registry
55        .send(move |client| {
56            client
57                .get(url.clone())
58                .header(reqwest::header::ACCEPT, accept.clone())
59        })
60        .await?;
61
62    let status = response.status();
63    if !status.is_success() {
64        return Err(BundleError::fetch(
65            &registry.name,
66            format!("manifest request failed: {status}"),
67        ));
68    }
69
70    let header_digest = response
71        .headers()
72        .get("docker-content-digest")
73        .and_then(|v| v.to_str().ok())
74        .map(str::to_owned);
75    let body = response
76        .bytes()
77        .await
78        .map_err(|e| BundleError::fetch(&registry.name, e))?;
79    let digest =
80        header_digest.unwrap_or_else(|| format!("sha256:{}", hex::encode(Sha256::digest(&body))));
81
82    let manifest: OciManifest = serde_json::from_slice(&body)
83        .map_err(|e| BundleError::fetch(&registry.name, format!("manifest does not parse: {e}")))?;
84    Ok((digest, manifest))
85}
86
87pub async fn pull_bundle_layer(
88    registry: &RegistryClient,
89    manifest: &OciManifest,
90    into: &Path,
91    max_bytes: u64,
92) -> BundleResult<FetchedBundle> {
93    let matching: Vec<&OciDescriptor> = manifest
94        .layers
95        .iter()
96        .filter(|l| l.media_type == BUNDLE_MEDIA_TYPE)
97        .collect();
98    let [layer] = matching.as_slice() else {
99        return Err(BundleError::fetch(
100            &registry.name,
101            format!(
102                "manifest carries {} layers of {BUNDLE_MEDIA_TYPE}, expected exactly one",
103                matching.len()
104            ),
105        ));
106    };
107
108    if layer.size > max_bytes {
109        return Err(BundleError::TooLarge { bytes: max_bytes });
110    }
111
112    let url = registry.url(&format!("/blobs/{}", layer.digest))?;
113    let response = registry.send(move |client| client.get(url.clone())).await?;
114    let status = response.status();
115    if !status.is_success() {
116        return Err(BundleError::fetch(
117            &registry.name,
118            format!("blob request failed: {status}"),
119        ));
120    }
121
122    let digest = stream_to_file(response, into, &registry.name, max_bytes).await?;
123    let expected = layer
124        .digest
125        .strip_prefix("sha256:")
126        .unwrap_or(&layer.digest);
127    if !digest.eq_ignore_ascii_case(expected) {
128        return Err(BundleError::fetch(
129            &registry.name,
130            format!(
131                "blob digest is sha256:{digest}, manifest declares {}",
132                layer.digest
133            ),
134        ));
135    }
136
137    Ok(FetchedBundle {
138        archive: into.to_path_buf(),
139        digest: format!("sha256:{digest}"),
140    })
141}