use anyhow::Result;
use dspy_rs::data::dataloader::DataLoader;
use rstest::rstest;
#[rstest]
#[cfg_attr(miri, ignore = "MIRI has issues with network operations")]
fn test_load_hf_awesome_chatgpt_prompts() -> Result<()> {
let input_keys = vec!["events".to_string(), "inputs".to_string()];
let output_keys = vec!["output".to_string()];
let examples = DataLoader::load_hf(
"zed-industries/zeta",
input_keys.clone(),
output_keys.clone(),
"", "train", true, )?;
assert!(
!examples.is_empty(),
"Should have loaded some examples from HuggingFace dataset"
);
let first_example = &examples[0];
assert_eq!(first_example.input_keys, input_keys);
assert_eq!(first_example.output_keys, output_keys);
let has_act = first_example.data.contains_key("act");
let has_prompt = first_example.data.contains_key("prompt");
assert!(
has_act || !first_example.keys().is_empty(),
"Example should contain 'act' field or have some data. Available fields: {:?}",
first_example.keys()
);
assert!(
has_prompt || !first_example.keys().is_empty(),
"Example should contain 'prompt' field or have some data. Available fields: {:?}",
first_example.keys()
);
if has_act && has_prompt {
let act_value = first_example.get("act", None);
let prompt_value = first_example.get("prompt", None);
assert!(!act_value.is_null(), "act field should not be null");
assert!(!prompt_value.is_null(), "prompt field should not be null");
let act_str = act_value.as_str().unwrap_or("");
let prompt_str = prompt_value.as_str().unwrap_or("");
assert!(!act_str.is_empty(), "act field should not be empty");
assert!(!prompt_str.is_empty(), "prompt field should not be empty");
}
Ok(())
}
#[rstest]
#[cfg_attr(miri, ignore = "MIRI has issues with network operations")]
fn test_load_csv_from_url() -> Result<()> {
let url = "https://people.sc.fsu.edu/~jburkardt/data/csv/snakes_count_10.csv";
let input_keys = vec!["Game Number".to_string()];
let output_keys = vec!["Game Length".to_string()];
let examples = DataLoader::load_csv(
url,
',', input_keys.clone(),
output_keys.clone(),
true, )?;
assert!(
!examples.is_empty(),
"Should have loaded some examples from CSV"
);
assert_eq!(
examples.len(),
10,
"Should have loaded exactly 10 game records"
);
let first_example = &examples[0];
assert_eq!(first_example.input_keys, input_keys);
assert_eq!(first_example.output_keys, output_keys);
assert!(
!first_example.data.is_empty(),
"Example should contain data"
);
Ok(())
}
#[rstest]
#[cfg_attr(miri, ignore = "MIRI has issues with network operations")]
fn test_load_json_from_url() -> Result<()> {
let url = "https://huggingface.co/xai-org/grok-2/raw/main/config.json";
let input_keys = vec!["vocab_size".to_string(), "hidden_size".to_string()];
let output_keys = vec![];
let examples = DataLoader::load_json(
url,
false, input_keys.clone(),
output_keys.clone(),
)?;
assert!(!examples.is_empty(), "Should have loaded data from JSON");
let config_example = &examples[0];
assert!(
config_example.data.contains_key("vocab_size"),
"Config should contain 'vocab_size' field"
);
assert!(
config_example.data.contains_key("hidden_size"),
"Config should contain 'hidden_size' field"
);
let vocab_size = config_example.get("vocab_size", None);
let hidden_size = config_example.get("hidden_size", None);
assert!(!vocab_size.is_null(), "vocab_size should not be null");
assert!(!hidden_size.is_null(), "hidden_size should not be null");
Ok(())
}
#[rstest]
#[cfg_attr(miri, ignore = "MIRI has issues with network operations")]
fn test_load_json_grok2_with_multiple_fields() -> Result<()> {
let url = "https://huggingface.co/xai-org/grok-2/raw/main/config.json";
let input_keys = vec![
"vocab_size".to_string(),
"hidden_size".to_string(),
"intermediate_size".to_string(),
"num_hidden_layers".to_string(),
];
let output_keys = vec![];
let examples = DataLoader::load_json(url, false, input_keys.clone(), output_keys.clone())?;
assert!(!examples.is_empty(), "Should have loaded data from JSON");
let config = &examples[0];
for key in &input_keys {
assert!(
config.data.contains_key(key),
"Config should contain '{key}' field"
);
let value = config.get(key, None);
assert!(!value.is_null(), "{key} should not be null");
}
Ok(())
}
#[rstest]
#[cfg_attr(miri, ignore = "MIRI has issues with network operations")]
fn test_load_csv_verify_columns() -> Result<()> {
let url = "https://people.sc.fsu.edu/~jburkardt/data/csv/snakes_count_10.csv";
let examples = DataLoader::load_csv(
url,
',',
vec![], vec![], true, )?;
assert!(!examples.is_empty(), "Should have loaded examples");
let first_example = &examples[0];
let keys = first_example.keys();
assert_eq!(examples.len(), 10, "Should have 10 game records");
for (i, example) in examples.iter().enumerate() {
assert_eq!(
example.keys().len(),
keys.len(),
"Row {i} should have same number of columns"
);
}
Ok(())
}
#[rstest]
#[cfg_attr(miri, ignore = "MIRI has issues with network operations")]
fn test_load_invalid_url_handling() {
let invalid_url = "https://invalid-url-that-does-not-exist.com/data.csv";
let result = DataLoader::load_csv(
invalid_url,
',',
vec!["col1".to_string()],
vec!["col2".to_string()],
true,
);
assert!(result.is_err(), "Should fail when loading from invalid URL");
}
#[rstest]
#[cfg_attr(miri, ignore = "MIRI has issues with network operations")]
fn test_load_hf_with_verbose() -> Result<()> {
let input_keys = vec!["events".to_string(), "inputs".to_string()];
let output_keys = vec!["output".to_string()];
let examples = DataLoader::load_hf(
"zed-industries/zeta",
input_keys.clone(),
output_keys.clone(),
"", "train", true, )?;
assert!(!examples.is_empty(), "Should have loaded examples");
for example in examples.iter().take(3) {
assert_eq!(example.input_keys, input_keys);
assert_eq!(example.output_keys, output_keys);
}
Ok(())
}