use crate::{
empty_body, Backend, Body, DirEntry, ExistsRequest, GetRequest, GetResponse, PutRequest,
PutResponse, StatRequest, StatResponse, DEFAULT_USER_AGENT, KEEP_ALIVE_INTERVAL,
POOL_MAX_IDLE_PER_HOST,
};
use async_trait::async_trait;
use dragonfly_api::common::v2::Range;
use dragonfly_client_config::dfdaemon::Config;
use dragonfly_client_core::{
error::{BackendError, ErrorType, OrErr},
Error, Result,
};
use dragonfly_client_util::{http::validate_ranged_response, tls::NoVerifier};
use futures::{StreamExt, TryStreamExt};
use reqwest::header::{HeaderMap, HeaderValue, AUTHORIZATION, CONTENT_LENGTH, RANGE, USER_AGENT};
use reqwest::Client;
use serde::Deserialize;
use std::error::Error as _;
use std::io::Error as IOError;
use std::sync::Arc;
use tokio_util::io::StreamReader;
use tracing::{debug, error, instrument};
use url::Url;
pub const SCHEME: &str = "opencsg";
const OPEN_CSG_BASE_URL: &str = "https://hub.opencsg.com/csg/";
#[derive(Default, Debug, Deserialize)]
#[serde(default)]
struct Repository {
siblings: Option<Vec<Sibling>>,
}
#[derive(Default, Debug, Deserialize)]
#[serde(default)]
struct Sibling {
rfilename: String,
size: Option<u64>,
lfs: Option<Lfs>,
r#type: Option<String>,
}
#[derive(Default, Debug, Deserialize)]
#[serde(default)]
struct Lfs {
size: Option<u64>,
}
#[derive(Debug, Clone)]
pub struct ParsedURL {
pub url: Url,
pub repository_id: String,
pub repository_type: RepositoryType,
pub file_path: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RepositoryType {
Model,
Dataset,
Space,
Code,
Mcp,
Skill,
}
impl RepositoryType {
pub fn as_str(&self) -> &'static str {
match self {
RepositoryType::Model => "models",
RepositoryType::Dataset => "datasets",
RepositoryType::Space => "spaces",
RepositoryType::Code => "codes",
RepositoryType::Mcp => "mcps",
RepositoryType::Skill => "skills",
}
}
}
impl TryFrom<Url> for ParsedURL {
type Error = Error;
fn try_from(url: Url) -> std::result::Result<Self, Self::Error> {
let host = url
.host_str()
.ok_or_else(|| Error::InvalidURI(url.to_string()))?;
let raw_path = format!("{}{}", host, url.path().trim_end_matches('/'));
let segments: Vec<&str> = raw_path.trim_matches('/').split('/').collect();
let (repository_type, offset) = match segments.first() {
Some(&"datasets") => (RepositoryType::Dataset, 1),
Some(&"spaces") => (RepositoryType::Space, 1),
Some(&"codes") => (RepositoryType::Code, 1),
Some(&"mcps") => (RepositoryType::Mcp, 1),
Some(&"skills") => (RepositoryType::Skill, 1),
Some(&"models") => (RepositoryType::Model, 1),
_ => (RepositoryType::Model, 0),
};
let remaining = &segments[offset..];
if remaining.len() < 2 {
return Err(Error::InvalidParameter);
}
let repository_id = format!("{}/{}", remaining[0], remaining[1]);
let file_path = if remaining.len() > 2 {
Some(remaining[2..].join("/"))
} else {
None
};
Ok(ParsedURL {
url,
repository_type,
repository_id,
file_path,
})
}
}
impl TryFrom<&str> for ParsedURL {
type Error = Error;
fn try_from(url: &str) -> std::result::Result<Self, Self::Error> {
let parsed_url = Url::parse(url).or_err(ErrorType::ParseError)?;
ParsedURL::try_from(parsed_url)
}
}
pub struct OpenCsg {
scheme: String,
client: Client,
}
impl OpenCsg {
pub fn new(config: Arc<Config>) -> Result<Self> {
let client_config_builder = rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(NoVerifier::new())
.with_no_client_auth();
let client = reqwest::Client::builder()
.no_gzip()
.no_brotli()
.no_zstd()
.no_deflate()
.hickory_dns(config.backend.enable_hickory_dns)
.use_preconfigured_tls(client_config_builder)
.pool_max_idle_per_host(POOL_MAX_IDLE_PER_HOST)
.tcp_keepalive(KEEP_ALIVE_INTERVAL)
.tcp_nodelay(true)
.build()?;
Ok(Self {
scheme: SCHEME.to_string(),
client,
})
}
fn resolve_base_urls(base_url: Option<&str>) -> Result<(Url, Url)> {
let base_url = Url::parse(base_url.unwrap_or(OPEN_CSG_BASE_URL))?;
let api_base_url = base_url.join("api/")?;
Ok((base_url, api_base_url))
}
fn build_download_url(
parsed_url: &ParsedURL,
file_path: &str,
revision: &str,
base_url: &Url,
) -> Result<Url> {
let path = match parsed_url.repository_type {
RepositoryType::Model => {
format!(
"{}/resolve/{}/{}",
parsed_url.repository_id, revision, file_path
)
}
RepositoryType::Dataset => {
format!(
"datasets/{}/resolve/{}/{}",
parsed_url.repository_id, revision, file_path
)
}
RepositoryType::Space => {
format!(
"spaces/{}/resolve/{}/{}",
parsed_url.repository_id, revision, file_path
)
}
RepositoryType::Code => {
format!(
"codes/{}/resolve/{}/{}",
parsed_url.repository_id, revision, file_path
)
}
RepositoryType::Mcp => {
format!(
"mcps/{}/resolve/{}/{}",
parsed_url.repository_id, revision, file_path
)
}
RepositoryType::Skill => {
format!(
"skills/{}/resolve/{}/{}",
parsed_url.repository_id, revision, file_path
)
}
};
Ok(base_url.join(&path)?)
}
fn build_repository_revision_url(
parsed_url: &ParsedURL,
revision: &str,
api_base_url: &Url,
) -> Result<Url> {
let path = match parsed_url.repository_type {
RepositoryType::Model => {
format!(
"models/{}/revision/{}?blobs=true",
parsed_url.repository_id, revision
)
}
RepositoryType::Dataset => {
format!(
"datasets/{}/revision/{}",
parsed_url.repository_id, revision
)
}
RepositoryType::Space => {
format!("spaces/{}/revision/{}", parsed_url.repository_id, revision)
}
RepositoryType::Code => {
format!("codes/{}/revision/{}", parsed_url.repository_id, revision)
}
RepositoryType::Mcp => {
format!("mcps/{}/revision/{}", parsed_url.repository_id, revision)
}
RepositoryType::Skill => {
format!("skills/{}/revision/{}", parsed_url.repository_id, revision)
}
};
Ok(api_base_url.join(&path)?)
}
fn build_opencsg_url(parsed_url: &ParsedURL, filename: &str) -> Result<Url> {
let url = match parsed_url.repository_type {
RepositoryType::Model => {
format!(
"{}://models/{}/{}",
SCHEME, parsed_url.repository_id, filename
)
}
RepositoryType::Dataset => {
format!(
"{}://datasets/{}/{}",
SCHEME, parsed_url.repository_id, filename
)
}
RepositoryType::Space => {
format!(
"{}://spaces/{}/{}",
SCHEME, parsed_url.repository_id, filename
)
}
RepositoryType::Code => {
format!(
"{}://codes/{}/{}",
SCHEME, parsed_url.repository_id, filename
)
}
RepositoryType::Mcp => {
format!(
"{}://mcps/{}/{}",
SCHEME, parsed_url.repository_id, filename
)
}
RepositoryType::Skill => {
format!(
"{}://skills/{}/{}",
SCHEME, parsed_url.repository_id, filename
)
}
};
Ok(Url::parse(&url)?)
}
fn build_request_headers(token: Option<String>, range: Option<Range>) -> Result<HeaderMap> {
let mut request_header = HeaderMap::new();
if let Some(range) = &range {
request_header.insert(
RANGE,
format!("bytes={}-{}", range.start, range.start + range.length - 1).parse()?,
);
};
request_header
.entry(USER_AGENT)
.or_insert(HeaderValue::from_static(DEFAULT_USER_AGENT));
if let Some(token) = token {
request_header.insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {token}")).unwrap(),
);
}
Ok(request_header)
}
}
#[async_trait]
impl Backend for OpenCsg {
fn scheme(&self) -> String {
self.scheme.clone()
}
#[instrument(skip_all)]
async fn stat(&self, request: StatRequest) -> Result<StatResponse> {
debug!(
"stat request {} {}: {:?}",
request.task_id, request.url, request.http_header
);
let request_header = Self::build_request_headers(
request.open_csg.as_ref().and_then(|csg| csg.token.clone()),
None,
)?;
let open_csg = request.open_csg.as_ref().ok_or_else(|| {
error!(
"stat request {} {}: missing OpenCSG information",
request.task_id, request.url
);
Error::InvalidParameter
})?;
let parsed_url = ParsedURL::try_from(request.url.as_str())?;
let (base_url, api_base_url) = Self::resolve_base_urls(open_csg.base_url.as_deref())?;
match &parsed_url.file_path {
Some(file_path) => {
let download_url = Self::build_download_url(
&parsed_url,
file_path,
&open_csg.revision,
&base_url,
)?;
let response = self
.client
.head(download_url.as_str())
.headers(request_header)
.timeout(request.timeout)
.send()
.await
.map_err(|err| {
error!(
"stat request failed {} {}: {}",
request.task_id, download_url, err
);
Error::BackendError(Box::new(BackendError {
message: err.to_string(),
status_code: None,
header: None,
}))
})?;
let response_status_code = response.status();
let response_header = response.headers().clone();
let content_length = match response_header.get(CONTENT_LENGTH) {
Some(content_length) => content_length.to_str()?.parse::<u64>().ok(),
None => response.content_length(),
};
if !response.status().is_success() {
error!(
"stat request failed {} {}: {}",
request.task_id, download_url, response_status_code
);
return Err(Error::BackendError(Box::new(BackendError {
message: response_status_code.to_string(),
status_code: Some(response_status_code),
header: Some(response_header),
})));
}
debug!(
"stat response {} {}: {:?} {:?} {:?}",
request.task_id,
download_url,
response_status_code,
content_length,
response_header
);
Ok(StatResponse {
success: response_status_code.is_success(),
content_length,
http_header: Some(response_header),
http_status_code: Some(response_status_code),
error_message: Some(response_status_code.to_string()),
entries: Vec::new(),
})
}
None => {
let repository_revision_url = Self::build_repository_revision_url(
&parsed_url,
&open_csg.revision,
&api_base_url,
)?;
let response = self
.client
.get(repository_revision_url.as_str())
.headers(request_header)
.timeout(request.timeout)
.send()
.await
.map_err(|err| {
error!(
"stat request failed {} {}: {}",
request.task_id, repository_revision_url, err
);
Error::BackendError(Box::new(BackendError {
message: err.to_string(),
status_code: None,
header: None,
}))
})?;
let response_status_code = response.status();
let response_header = response.headers().clone();
let content_length = match response_header.get(CONTENT_LENGTH) {
Some(content_length) => content_length.to_str()?.parse::<u64>().ok(),
None => response.content_length(),
};
if !response.status().is_success() {
error!(
"stat request failed {} {}: {}",
request.task_id, repository_revision_url, response_status_code
);
return Err(Error::BackendError(Box::new(BackendError {
message: response_status_code.to_string(),
status_code: Some(response_status_code),
header: Some(response_header),
})));
}
let text = response.text().await.map_err(|err| {
error!(
"stat request failed {} {}: {}",
request.task_id, repository_revision_url, err
);
Error::BackendError(Box::new(BackendError {
message: err.to_string(),
status_code: None,
header: None,
}))
})?;
let repository: Repository = serde_json::from_str(&text).map_err(|err| {
error!(
"stat request failed {} {}: {}",
request.task_id, repository_revision_url, err
);
Error::BackendError(Box::new(BackendError {
message: err.to_string(),
status_code: None,
header: None,
}))
})?;
let entries: Vec<DirEntry> = repository
.siblings
.unwrap_or_default()
.into_iter()
.filter(|sibling: &Sibling| sibling.r#type.as_deref() != Some("tree"))
.map(|sibling: Sibling| -> Result<DirEntry> {
let opencsg_url: Url =
Self::build_opencsg_url(&parsed_url, &sibling.rfilename)?;
let content_length: u64 = sibling
.lfs
.and_then(|lfs: Lfs| lfs.size)
.or(sibling.size)
.unwrap_or(0);
Ok(DirEntry {
url: opencsg_url.to_string(),
content_length: content_length as usize,
is_dir: false,
})
})
.collect::<Result<Vec<_>>>()?;
debug!(
"stat response {} {}: {:?} {:?} {:?}",
request.task_id,
repository_revision_url,
response_status_code,
content_length,
response_header
);
Ok(StatResponse {
success: response_status_code.is_success(),
content_length,
http_header: Some(response_header),
http_status_code: Some(response_status_code),
error_message: Some(response_status_code.to_string()),
entries,
})
}
}
}
#[instrument(skip_all)]
async fn get(&self, request: GetRequest) -> Result<GetResponse<Body>> {
debug!(
"get request {} {} {}: {:?}",
request.task_id, request.piece_id, request.url, request.http_header
);
let request_header = Self::build_request_headers(
request.open_csg.as_ref().and_then(|csg| csg.token.clone()),
request.range,
)?;
let open_csg = request.open_csg.as_ref().ok_or_else(|| {
error!(
"get request {} {}: missing OpenCSG information",
request.task_id, request.url
);
Error::InvalidParameter
})?;
let parsed_url = ParsedURL::try_from(request.url.as_str())?;
let Some(file_path) = &parsed_url.file_path else {
error!(
"get request {} {}: URL must specify a file path",
request.task_id, request.url
);
return Err(Error::InvalidParameter);
};
let (base_url, _) = Self::resolve_base_urls(open_csg.base_url.as_deref())?;
let download_url =
Self::build_download_url(&parsed_url, file_path, &open_csg.revision, &base_url)?;
let response = match self
.client
.get(download_url.as_str())
.headers(request_header)
.timeout(request.timeout)
.send()
.await
{
Ok(response) => response,
Err(err) => {
error!(
"get request failed {} {} {}: {}",
request.task_id, request.piece_id, download_url, err
);
return Ok(GetResponse {
success: false,
http_header: None,
http_status_code: None,
reader: empty_body(),
error_message: Some(err.to_string()),
});
}
};
let response_header = response.headers().clone();
let response_status_code = response.status();
if let Err(err) =
validate_ranged_response(request.range, response_status_code, &response_header)
{
error!(
"get request failed {} {} {}: {}",
request.task_id, request.piece_id, download_url, err
);
return Ok(GetResponse {
success: false,
http_header: Some(response_header),
http_status_code: Some(response_status_code),
reader: empty_body(),
error_message: Some(err.to_string()),
});
}
let response_reader = StreamReader::new(
response
.bytes_stream()
.map_err(move |err| {
let mut chain = err.to_string();
let mut source = err.source();
while let Some(err) = source {
chain.push_str(": ");
chain.push_str(&err.to_string());
source = err.source();
}
IOError::other(chain)
})
.boxed(),
);
debug!(
"get response {} {}: {:?} {:?}",
request.task_id, request.piece_id, response_status_code, response_header,
);
Ok(GetResponse {
success: response_status_code.is_success(),
http_header: Some(response_header),
http_status_code: Some(response_status_code),
reader: response_reader,
error_message: Some(response_status_code.to_string()),
})
}
async fn put(&self, _request: PutRequest) -> Result<PutResponse> {
unimplemented!()
}
#[instrument(skip_all)]
async fn exists(&self, request: ExistsRequest) -> Result<bool> {
debug!(
"exists request {} {}: {:?}",
request.task_id, request.url, request.http_header
);
let request_header = Self::build_request_headers(
request.open_csg.as_ref().and_then(|csg| csg.token.clone()),
None,
)?;
let open_csg = request.open_csg.as_ref().ok_or_else(|| {
error!(
"exists request {} {}: missing OpenCSG information",
request.task_id, request.url
);
Error::InvalidParameter
})?;
let parsed_url = ParsedURL::try_from(request.url.as_str())?;
let (base_url, api_base_url) = Self::resolve_base_urls(open_csg.base_url.as_deref())?;
match &parsed_url.file_path {
Some(file_path) => {
let download_url = Self::build_download_url(
&parsed_url,
file_path,
&open_csg.revision,
&base_url,
)?;
let response = self
.client
.head(download_url.as_str())
.headers(request_header)
.timeout(request.timeout)
.send()
.await
.inspect_err(|err| {
error!(
"exists request failed {} {}: {}",
request.task_id, request.url, err
);
})?;
let response_status_code = response.status();
debug!(
"exists response {} {}: {:?} {:?}",
request.task_id,
request.url,
response_status_code,
response.headers()
);
Ok(response_status_code.is_success())
}
None => {
let repository_revision_url = Self::build_repository_revision_url(
&parsed_url,
&open_csg.revision,
&api_base_url,
)?;
let response = self
.client
.head(repository_revision_url.as_str())
.headers(request_header)
.timeout(request.timeout)
.send()
.await
.inspect_err(|err| {
error!(
"exists request failed {} {}: {}",
request.task_id, request.url, err
);
})?;
let response_status_code = response.status();
debug!(
"exists response {} {}: {:?} {:?}",
request.task_id,
request.url,
response_status_code,
response.headers()
);
Ok(response_status_code.is_success())
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::DEFAULT_USER_AGENT;
use dragonfly_api::common::v2::OpenCsg as OpenCsgOptions;
use std::time::Duration;
use wiremock::{
matchers::{header, method, path, query_param},
Mock, MockServer, ResponseTemplate,
};
#[test]
fn test_parse_url_simple() {
let parsed_url = ParsedURL::try_from("opencsg://OpenCSG/csg-wukong-1B").unwrap();
assert_eq!(parsed_url.repository_id, "OpenCSG/csg-wukong-1B");
assert_eq!(parsed_url.repository_type, RepositoryType::Model);
assert!(parsed_url.file_path.is_none());
}
#[test]
fn test_parse_url_with_file() {
let parsed_url =
ParsedURL::try_from("opencsg://OpenCSG/csg-wukong-1B/model.safetensors").unwrap();
assert_eq!(parsed_url.repository_id, "OpenCSG/csg-wukong-1B");
assert_eq!(parsed_url.repository_type, RepositoryType::Model);
assert_eq!(parsed_url.file_path, Some("model.safetensors".to_string()));
}
#[test]
fn test_parse_url_with_nested_path() {
let parsed_url =
ParsedURL::try_from("opencsg://OpenCSG/csg-wukong-1B/models/v1/model.bin").unwrap();
assert_eq!(parsed_url.repository_id, "OpenCSG/csg-wukong-1B");
assert_eq!(parsed_url.repository_type, RepositoryType::Model);
assert_eq!(
parsed_url.file_path,
Some("models/v1/model.bin".to_string())
);
}
#[test]
fn test_parse_url_dataset() {
let parsed_url =
ParsedURL::try_from("opencsg://datasets/OpenCSG/chinese-fineweb-edu").unwrap();
assert_eq!(parsed_url.repository_id, "OpenCSG/chinese-fineweb-edu");
assert_eq!(parsed_url.repository_type, RepositoryType::Dataset);
assert!(parsed_url.file_path.is_none());
}
#[test]
fn test_parse_url_dataset_with_path() {
let parsed_url =
ParsedURL::try_from("opencsg://datasets/OpenCSG/chinese-fineweb-edu/train.json")
.unwrap();
assert_eq!(parsed_url.repository_id, "OpenCSG/chinese-fineweb-edu");
assert_eq!(parsed_url.repository_type, RepositoryType::Dataset);
assert_eq!(parsed_url.file_path, Some("train.json".to_string()));
}
#[test]
fn test_parse_url_space() {
let parsed_url = ParsedURL::try_from("opencsg://spaces/owner/repo").unwrap();
assert_eq!(parsed_url.repository_id, "owner/repo");
assert_eq!(parsed_url.repository_type, RepositoryType::Space);
assert!(parsed_url.file_path.is_none());
}
#[test]
fn test_parse_url_code() {
let parsed_url = ParsedURL::try_from("opencsg://codes/owner/repo").unwrap();
assert_eq!(parsed_url.repository_id, "owner/repo");
assert_eq!(parsed_url.repository_type, RepositoryType::Code);
assert!(parsed_url.file_path.is_none());
}
#[test]
fn test_parse_url_mcp() {
let parsed_url = ParsedURL::try_from("opencsg://mcps/owner/repo").unwrap();
assert_eq!(parsed_url.repository_id, "owner/repo");
assert_eq!(parsed_url.repository_type, RepositoryType::Mcp);
assert!(parsed_url.file_path.is_none());
}
#[test]
fn test_parse_url_skill() {
let parsed_url = ParsedURL::try_from("opencsg://skills/owner/repo").unwrap();
assert_eq!(parsed_url.repository_id, "owner/repo");
assert_eq!(parsed_url.repository_type, RepositoryType::Skill);
assert!(parsed_url.file_path.is_none());
}
#[test]
fn test_parse_url_explicit_model_type() {
let parsed_url =
ParsedURL::try_from("opencsg://models/OpenCSG/csg-wukong-1B/model.safetensors")
.unwrap();
assert_eq!(parsed_url.repository_id, "OpenCSG/csg-wukong-1B");
assert_eq!(parsed_url.repository_type, RepositoryType::Model);
assert_eq!(parsed_url.file_path, Some("model.safetensors".to_string()));
}
#[test]
fn test_parse_url_missing_repo() {
let result = ParsedURL::try_from("opencsg://owner");
assert!(result.is_err());
}
#[test]
fn test_build_download_url_model() {
let parsed_url =
ParsedURL::try_from("opencsg://OpenCSG/csg-wukong-1B/model.safetensors").unwrap();
let url = OpenCsg::build_download_url(
&parsed_url,
"model.safetensors",
"main",
&Url::parse(OPEN_CSG_BASE_URL).unwrap(),
)
.unwrap();
assert_eq!(
url.as_str(),
"https://hub.opencsg.com/csg/OpenCSG/csg-wukong-1B/resolve/main/model.safetensors"
);
}
#[test]
fn test_build_download_url_dataset() {
let parsed_url =
ParsedURL::try_from("opencsg://datasets/OpenCSG/chinese-fineweb-edu/train.json")
.unwrap();
let url = OpenCsg::build_download_url(
&parsed_url,
"train.json",
"main",
&Url::parse(OPEN_CSG_BASE_URL).unwrap(),
)
.unwrap();
assert_eq!(
url.as_str(),
"https://hub.opencsg.com/csg/datasets/OpenCSG/chinese-fineweb-edu/resolve/main/train.json"
);
}
#[test]
fn test_build_repository_revision_url_model() {
let parsed_url = ParsedURL::try_from("opencsg://OpenCSG/csg-wukong-1B").unwrap();
let url = OpenCsg::build_repository_revision_url(
&parsed_url,
"main",
&Url::parse("https://hub.opencsg.com/csg/api/").unwrap(),
)
.unwrap();
assert_eq!(
url.as_str(),
"https://hub.opencsg.com/csg/api/models/OpenCSG/csg-wukong-1B/revision/main?blobs=true"
);
}
#[test]
fn test_build_repository_revision_url_dataset() {
let parsed_url =
ParsedURL::try_from("opencsg://datasets/OpenCSG/chinese-fineweb-edu").unwrap();
let url = OpenCsg::build_repository_revision_url(
&parsed_url,
"main",
&Url::parse("https://hub.opencsg.com/csg/api/").unwrap(),
)
.unwrap();
assert_eq!(
url.as_str(),
"https://hub.opencsg.com/csg/api/datasets/OpenCSG/chinese-fineweb-edu/revision/main"
);
}
#[test]
fn test_build_opencsg_url_model() {
let parsed_url = ParsedURL::try_from("opencsg://OpenCSG/csg-wukong-1B").unwrap();
let url = OpenCsg::build_opencsg_url(&parsed_url, "model.safetensors").unwrap();
assert_eq!(
url.as_str(),
"opencsg://models/OpenCSG/csg-wukong-1B/model.safetensors"
);
}
#[test]
fn test_build_opencsg_url_dataset() {
let parsed_url =
ParsedURL::try_from("opencsg://datasets/OpenCSG/chinese-fineweb-edu").unwrap();
let url = OpenCsg::build_opencsg_url(&parsed_url, "train.json").unwrap();
assert_eq!(
url.as_str(),
"opencsg://datasets/OpenCSG/chinese-fineweb-edu/train.json"
);
}
#[test]
fn test_resolve_base_urls() {
let (base_url, api_base_url) =
OpenCsg::resolve_base_urls(Some("https://hub-mirror.example.com/csg/")).unwrap();
assert_eq!(base_url.as_str(), "https://hub-mirror.example.com/csg/");
assert_eq!(
api_base_url.as_str(),
"https://hub-mirror.example.com/csg/api/"
);
}
#[test]
fn test_build_headers_default_user_agent() {
let request_header = OpenCsg::build_request_headers(None, None).unwrap();
assert_eq!(
request_header.get(USER_AGENT).unwrap(),
HeaderValue::from_static(DEFAULT_USER_AGENT)
);
}
#[test]
fn test_build_headers_preserves_request_headers() {
let request_headers =
OpenCsg::build_request_headers(Some("test-token".to_string()), None).unwrap();
assert_eq!(
request_headers.get(reqwest::header::AUTHORIZATION).unwrap(),
"Bearer test-token"
);
assert_eq!(
request_headers.get(USER_AGENT).unwrap(),
HeaderValue::from_static(DEFAULT_USER_AGENT)
);
}
#[test]
fn test_build_headers_with_range() {
let request_headers = OpenCsg::build_request_headers(
None,
Some(Range {
start: 0,
length: 1024,
}),
)
.unwrap();
assert_eq!(
request_headers.get(RANGE).unwrap(),
HeaderValue::from_static("bytes=0-1023")
);
}
#[test]
fn test_parse_repository_null_siblings() {
let repository: Repository = serde_json::from_str(r#"{"siblings":null}"#).unwrap();
assert!(repository.siblings.unwrap_or_default().is_empty());
}
#[tokio::test]
async fn test_stat_repository() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/api/models/owner/repo/revision/main"))
.and(query_param("blobs", "true"))
.and(header("authorization", "Bearer secret"))
.and(header("user-agent", DEFAULT_USER_AGENT))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"siblings": [
{"rfilename": "nested/config file.json", "size": 12},
{"rfilename": "model.bin", "size": 128, "lfs": {"size": 4096}},
{"rfilename": "README.md"},
{"rfilename": "ignored", "type": "tree"}
]
})))
.mount(&server)
.await;
let backend = OpenCsg::new(Arc::new(Config::default())).unwrap();
let response = backend
.stat(StatRequest {
task_id: "task".to_string(),
url: "opencsg://owner/repo".to_string(),
http_header: None,
timeout: Duration::from_secs(5),
client_cert: None,
object_storage: None,
hdfs: None,
hugging_face: None,
model_scope: None,
open_csg: Some(OpenCsgOptions {
revision: "main".to_string(),
token: Some("secret".to_string()),
base_url: Some(server.uri()),
}),
})
.await
.unwrap();
assert!(response.success);
assert_eq!(response.entries.len(), 3);
assert_eq!(
response.entries[0],
DirEntry {
url: "opencsg://models/owner/repo/nested/config%20file.json".to_string(),
content_length: 12,
is_dir: false,
}
);
assert_eq!(response.entries[1].content_length, 4096);
assert_eq!(response.entries[2].content_length, 0);
}
#[tokio::test]
async fn test_stat_dataset_repository() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/api/datasets/owner/repo/revision/main"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"siblings": [
{"rfilename": "nested/train.json"},
{"rfilename": "README.md"}
]
})))
.mount(&server)
.await;
let backend = OpenCsg::new(Arc::new(Config::default())).unwrap();
let response = backend
.stat(StatRequest {
task_id: "task".to_string(),
url: "opencsg://datasets/owner/repo".to_string(),
http_header: None,
timeout: Duration::from_secs(5),
client_cert: None,
object_storage: None,
hdfs: None,
hugging_face: None,
model_scope: None,
open_csg: Some(OpenCsgOptions {
revision: "main".to_string(),
token: None,
base_url: Some(server.uri()),
}),
})
.await
.unwrap();
assert!(response.success);
assert_eq!(response.entries.len(), 2);
assert_eq!(
response.entries[0],
DirEntry {
url: "opencsg://datasets/owner/repo/nested/train.json".to_string(),
content_length: 0,
is_dir: false,
}
);
}
#[tokio::test]
async fn test_stat_file() {
let server = MockServer::start().await;
Mock::given(method("HEAD"))
.and(path("/owner/repo/resolve/main/model.bin"))
.respond_with(ResponseTemplate::new(200).insert_header("content-length", "4096"))
.mount(&server)
.await;
let backend = OpenCsg::new(Arc::new(Config::default())).unwrap();
let response = backend
.stat(StatRequest {
task_id: "task".to_string(),
url: "opencsg://owner/repo/model.bin".to_string(),
http_header: None,
timeout: Duration::from_secs(5),
client_cert: None,
object_storage: None,
hdfs: None,
hugging_face: None,
model_scope: None,
open_csg: Some(OpenCsgOptions {
revision: "main".to_string(),
token: None,
base_url: Some(server.uri()),
}),
})
.await
.unwrap();
assert!(response.success);
assert_eq!(response.content_length, Some(4096));
}
#[tokio::test]
async fn test_stat_file_with_error_status() {
let server = MockServer::start().await;
Mock::given(method("HEAD"))
.and(path("/monkey/Qwen/resolve/main/Qwen3.5-0.8B"))
.respond_with(ResponseTemplate::new(404))
.mount(&server)
.await;
let backend = OpenCsg::new(Arc::new(Config::default())).unwrap();
let err = backend
.stat(StatRequest {
task_id: "task".to_string(),
url: "opencsg://monkey/Qwen/Qwen3.5-0.8B".to_string(),
http_header: None,
timeout: Duration::from_secs(5),
client_cert: None,
object_storage: None,
hdfs: None,
hugging_face: None,
model_scope: None,
open_csg: Some(OpenCsgOptions {
revision: "main".to_string(),
token: None,
base_url: Some(server.uri()),
}),
})
.await
.unwrap_err();
assert!(matches!(
err,
Error::BackendError(err) if err.status_code == Some(reqwest::StatusCode::NOT_FOUND)
));
}
#[tokio::test]
async fn test_stat_repository_with_error_status() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/api/models/owner/repo/revision/main"))
.respond_with(ResponseTemplate::new(401))
.mount(&server)
.await;
let backend = OpenCsg::new(Arc::new(Config::default())).unwrap();
let err = backend
.stat(StatRequest {
task_id: "task".to_string(),
url: "opencsg://owner/repo".to_string(),
http_header: None,
timeout: Duration::from_secs(5),
client_cert: None,
object_storage: None,
hdfs: None,
hugging_face: None,
model_scope: None,
open_csg: Some(OpenCsgOptions {
revision: "main".to_string(),
token: None,
base_url: Some(server.uri()),
}),
})
.await
.unwrap_err();
assert!(matches!(
err,
Error::BackendError(err) if err.status_code == Some(reqwest::StatusCode::UNAUTHORIZED)
));
}
#[tokio::test]
async fn test_get_propagates_range_header() {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/owner/repo/resolve/main/model.bin"))
.and(header("range", "bytes=10-29"))
.respond_with(
ResponseTemplate::new(206)
.insert_header("content-range", "bytes 10-29/100")
.set_body_string("partial content here"),
)
.mount(&server)
.await;
let backend = OpenCsg::new(Arc::new(Config::default())).unwrap();
let mut response = backend
.get(GetRequest {
task_id: "task".to_string(),
piece_id: "piece".to_string(),
url: "opencsg://owner/repo/model.bin".to_string(),
range: Some(Range {
start: 10,
length: 20,
}),
http_header: None,
timeout: Duration::from_secs(5),
client_cert: None,
object_storage: None,
hdfs: None,
hugging_face: None,
model_scope: None,
open_csg: Some(OpenCsgOptions {
revision: "main".to_string(),
token: None,
base_url: Some(server.uri()),
}),
})
.await
.unwrap();
assert!(response.success);
assert_eq!(response.text().await.unwrap(), "partial content here");
}
#[tokio::test]
async fn test_get_follows_lfs_redirect() {
let server = MockServer::start().await;
let object_server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/owner/repo/resolve/main/model.bin"))
.and(header("range", "bytes=10-29"))
.respond_with(ResponseTemplate::new(302).insert_header(
"location",
format!("{}/objects/model.bin", object_server.uri()),
))
.mount(&server)
.await;
Mock::given(method("GET"))
.and(path("/objects/model.bin"))
.and(header("range", "bytes=10-29"))
.respond_with(
ResponseTemplate::new(206)
.insert_header("content-range", "bytes 10-29/100")
.set_body_string("redirected lfs data!"),
)
.mount(&object_server)
.await;
let backend = OpenCsg::new(Arc::new(Config::default())).unwrap();
let mut response = backend
.get(GetRequest {
task_id: "task".to_string(),
piece_id: "piece".to_string(),
url: "opencsg://owner/repo/model.bin".to_string(),
range: Some(Range {
start: 10,
length: 20,
}),
http_header: None,
timeout: Duration::from_secs(5),
client_cert: None,
object_storage: None,
hdfs: None,
hugging_face: None,
model_scope: None,
open_csg: Some(OpenCsgOptions {
revision: "main".to_string(),
token: Some("secret".to_string()),
base_url: Some(server.uri()),
}),
})
.await
.unwrap();
assert!(response.success);
assert_eq!(response.text().await.unwrap(), "redirected lfs data!");
let object_requests = object_server.received_requests().await.unwrap();
assert_eq!(object_requests.len(), 1);
assert!(object_requests[0].headers.get("authorization").is_none());
}
#[tokio::test]
async fn test_exists() {
let server = MockServer::start().await;
Mock::given(method("HEAD"))
.and(path("/owner/repo/resolve/main/model.bin"))
.respond_with(ResponseTemplate::new(200))
.mount(&server)
.await;
Mock::given(method("HEAD"))
.and(path("/api/models/owner/repo/revision/main"))
.and(query_param("blobs", "true"))
.respond_with(ResponseTemplate::new(200))
.mount(&server)
.await;
let backend = OpenCsg::new(Arc::new(Config::default())).unwrap();
for (url, expected) in [
("opencsg://owner/repo/model.bin", true),
("opencsg://owner/repo", true),
("opencsg://owner/repo/missing.bin", false),
] {
let exists = backend
.exists(ExistsRequest {
task_id: "task".to_string(),
url: url.to_string(),
http_header: None,
timeout: Duration::from_secs(5),
client_cert: None,
object_storage: None,
hdfs: None,
hugging_face: None,
model_scope: None,
open_csg: Some(OpenCsgOptions {
revision: "main".to_string(),
token: None,
base_url: Some(server.uri()),
}),
})
.await
.unwrap();
assert_eq!(exists, expected);
}
}
}