osdk_core/model/provider/
modelscope.rs1use 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}