Skip to main content

osdk_core/model/provider/
modelscope.rs

1use async_trait::async_trait;
2use serde::Deserialize;
3
4use crate::backend::Ctx;
5use crate::error::{Error, Result};
6use crate::model::provider::{get_cached_json, ModelProvider, RemoteModelFile, RemoteSnapshot};
7use crate::model::{ModelRef, ProviderId};
8
9pub struct ModelScope {
10    token: Option<String>,
11    allow_auth: bool,
12}
13
14impl ModelScope {
15    pub fn new(allow_auth: bool) -> Self {
16        Self {
17            token: None,
18            allow_auth,
19        }
20    }
21
22    #[cfg(test)]
23    pub fn with_token(token: impl Into<String>) -> Self {
24        Self {
25            token: Some(token.into()),
26            allow_auth: true,
27        }
28    }
29}
30
31impl Default for ModelScope {
32    fn default() -> Self {
33        Self::new(true)
34    }
35}
36
37#[derive(Debug, Deserialize)]
38struct ApiResponse<T> {
39    #[serde(rename = "Code")]
40    code: i64,
41    #[serde(rename = "Success", default)]
42    success: bool,
43    #[serde(rename = "Message", default)]
44    message: String,
45    #[serde(rename = "Data")]
46    data: T,
47}
48
49#[derive(Debug, Deserialize)]
50struct FilesData {
51    #[serde(rename = "Files", default)]
52    files: Vec<FileInfo>,
53}
54
55#[derive(Debug, Deserialize)]
56struct FileInfo {
57    #[serde(rename = "Path")]
58    path: String,
59    #[serde(rename = "Size")]
60    size: u64,
61    #[serde(rename = "Sha256")]
62    sha256: String,
63    #[serde(rename = "Type", default)]
64    file_type: String,
65}
66
67#[async_trait]
68impl ModelProvider for ModelScope {
69    async fn resolve(
70        &self,
71        ctx: &Ctx,
72        reference: &ModelRef,
73        endpoint: &str,
74    ) -> Result<RemoteSnapshot> {
75        if reference.provider != ProviderId::ModelScope {
76            return Err(Error::config(format!(
77                "ModelScope provider cannot resolve {}",
78                reference.provider
79            )));
80        }
81        let endpoint = endpoint.trim_end_matches('/');
82        let metadata_url = files_url(endpoint, &reference.repository, &reference.revision)?;
83        let headers = if self.allow_auth {
84            auth_headers(self.token.as_deref())
85        } else {
86            Vec::new()
87        };
88        let cache_identity = format!(
89            "{}:{}:{}@{}",
90            reference.provider, endpoint, reference.repository, reference.revision
91        );
92        let response: ApiResponse<FilesData> = get_cached_json(
93            ctx,
94            reference.provider.as_str(),
95            &cache_identity,
96            &metadata_url,
97            &headers,
98        )
99        .await?;
100        if !response.success || response.code != 200 {
101            return Err(Error::other(format!(
102                "ModelScope API failed with code {}: {}",
103                response.code, response.message
104            )));
105        }
106
107        let mut files = Vec::new();
108        for file in response.data.files {
109            if file.file_type.eq_ignore_ascii_case("tree") {
110                continue;
111            }
112            crate::model::safe_relative_path(&file.path)?;
113            if file.sha256.len() != 64
114                || !file
115                    .sha256
116                    .chars()
117                    .all(|character| character.is_ascii_hexdigit())
118            {
119                return Err(Error::other(format!(
120                    "ModelScope file {} has no valid SHA-256",
121                    file.path
122                )));
123            }
124            files.push(RemoteModelFile {
125                url: download_url(
126                    endpoint,
127                    &reference.repository,
128                    &reference.revision,
129                    &file.path,
130                )?,
131                path: file.path,
132                size: Some(file.size),
133                etag: Some(file.sha256.clone()),
134                sha256: Some(file.sha256),
135                headers: headers.clone(),
136            });
137        }
138        if files.is_empty() {
139            return Err(Error::other(format!(
140                "ModelScope repository {}@{} contains no files",
141                reference.repository, reference.revision
142            )));
143        }
144        files.sort_by(|left, right| left.path.cmp(&right.path));
145        let revision = manifest_revision(&reference.revision, &files);
146        Ok(RemoteSnapshot {
147            revision,
148            endpoint: endpoint.to_string(),
149            files,
150        })
151    }
152}
153
154fn files_url(endpoint: &str, repository: &str, revision: &str) -> Result<String> {
155    let mut url = repo_url(endpoint, repository)?;
156    {
157        let mut segments = url
158            .path_segments_mut()
159            .map_err(|_| Error::config("ModelScope endpoint cannot be a base URL"))?;
160        segments.extend(["repo", "files"]);
161    }
162    url.query_pairs_mut()
163        .append_pair("Revision", revision)
164        .append_pair("Recursive", "true");
165    Ok(url.into())
166}
167
168fn download_url(
169    endpoint: &str,
170    repository: &str,
171    revision: &str,
172    file_path: &str,
173) -> Result<String> {
174    let mut url = repo_url(endpoint, repository)?;
175    url.path_segments_mut()
176        .map_err(|_| Error::config("ModelScope endpoint cannot be a base URL"))?
177        .push("repo");
178    url.query_pairs_mut()
179        .append_pair("Revision", revision)
180        .append_pair("FilePath", file_path);
181    Ok(url.into())
182}
183
184fn repo_url(endpoint: &str, repository: &str) -> Result<reqwest::Url> {
185    let mut url = reqwest::Url::parse(endpoint)
186        .map_err(|error| Error::config(format!("invalid ModelScope endpoint: {error}")))?;
187    {
188        let mut segments = url
189            .path_segments_mut()
190            .map_err(|_| Error::config("ModelScope endpoint cannot be a base URL"))?;
191        segments.pop_if_empty();
192        segments.extend(["api", "v1", "models"]);
193        for part in repository.split('/') {
194            segments.push(part);
195        }
196    }
197    Ok(url)
198}
199
200fn manifest_revision(requested: &str, files: &[RemoteModelFile]) -> String {
201    let mut hasher = blake3::Hasher::new();
202    for file in files {
203        hasher.update(file.path.as_bytes());
204        hasher.update(b"\0");
205        hasher.update(file.size.unwrap_or_default().to_string().as_bytes());
206        hasher.update(b"\0");
207        hasher.update(file.sha256.as_deref().unwrap_or_default().as_bytes());
208        hasher.update(b"\0");
209    }
210    format!("{requested}+manifest-{}", &hasher.finalize().to_hex()[..16])
211}
212
213fn auth_headers(explicit: Option<&str>) -> Vec<(String, String)> {
214    let token = explicit.map(str::to_string).or_else(|| {
215        ["OSDK_MODELSCOPE_TOKEN", "MODELSCOPE_API_TOKEN"]
216            .iter()
217            .find_map(|key| {
218                std::env::var(key).ok().and_then(|value| {
219                    let value = value.trim().to_string();
220                    (!value.is_empty()).then_some(value)
221                })
222            })
223    });
224    match token {
225        Some(token) => vec![
226            ("Authorization".into(), format!("Bearer {token}")),
227            ("Cookie".into(), format!("m_session_id={token}")),
228        ],
229        None => Vec::new(),
230    }
231}
232
233#[cfg(test)]
234mod tests {
235    use super::*;
236
237    #[test]
238    fn builds_modelscope_urls_and_stable_manifest_revision() {
239        let files = files_url(
240            "https://modelscope.example.test",
241            "owner/repo",
242            "release/v1",
243        )
244        .unwrap();
245        assert_eq!(
246            files,
247            "https://modelscope.example.test/api/v1/models/owner/repo/repo/files?Revision=release%2Fv1&Recursive=true"
248        );
249        let download = download_url(
250            "https://modelscope.example.test",
251            "owner/repo",
252            "master",
253            "weights/model.safetensors",
254        )
255        .unwrap();
256        assert_eq!(
257            download,
258            "https://modelscope.example.test/api/v1/models/owner/repo/repo?Revision=master&FilePath=weights%2Fmodel.safetensors"
259        );
260        let files = vec![RemoteModelFile {
261            path: "config.json".into(),
262            size: Some(10),
263            sha256: Some("a".repeat(64)),
264            etag: None,
265            url: String::new(),
266            headers: Vec::new(),
267        }];
268        assert_eq!(
269            manifest_revision("master", &files),
270            manifest_revision("master", &files)
271        );
272        assert!(manifest_revision("master", &files).starts_with("master+manifest-"));
273    }
274}