From 4693c7373c5f8d10ee8af732aad12d901602c65c Mon Sep 17 00:00:00 2001 From: Bryan Helmkamp Date: Wed, 4 Mar 2026 00:14:15 -0500 Subject: [PATCH] Add 49 unit tests for uncovered pure logic in arc-llm Cover cli.rs formatting/parsing/resolve/apply_options, common.rs parse_error_body/extract_system_prompt/parse_retry_after, tools.rs args_type_name, provider.rs validate_tool_choice, and types.rs ToolChoice::mode_str. Co-Authored-By: Claude Opus 4.6 --- crates/arc-llm/src/cli.rs | 193 +++++++++++++++++++++++++ crates/arc-llm/src/provider.rs | 68 +++++++++ crates/arc-llm/src/providers/common.rs | 128 ++++++++++++++++ crates/arc-llm/src/tools.rs | 32 ++++ crates/arc-llm/src/types.rs | 20 +++ 5 files changed, 441 insertions(+) diff --git a/crates/arc-llm/src/cli.rs b/crates/arc-llm/src/cli.rs index 8c54f4bca..50dad795a 100644 --- a/crates/arc-llm/src/cli.rs +++ b/crates/arc-llm/src/cli.rs @@ -442,3 +442,196 @@ async fn test_models(provider: Option<&str>, model: Option<&str>, s: &Styles) -> Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + + // --- parse_option --- + + #[test] + fn parse_option_valid() { + let (k, v) = parse_option("temperature=0.7").unwrap(); + assert_eq!(k, "temperature"); + assert_eq!(v, "0.7"); + } + + #[test] + fn parse_option_value_with_equals() { + let (k, v) = parse_option("key=a=b").unwrap(); + assert_eq!(k, "key"); + assert_eq!(v, "a=b"); + } + + #[test] + fn parse_option_no_equals() { + assert!(parse_option("nope").is_err()); + } + + // --- format_context_window --- + + #[test] + fn format_context_window_millions() { + assert_eq!(format_context_window(1_000_000), "1m"); + } + + #[test] + fn format_context_window_thousands() { + assert_eq!(format_context_window(128_000), "128k"); + } + + #[test] + fn format_context_window_small() { + assert_eq!(format_context_window(400), "400"); + } + + #[test] + fn format_context_window_rounds_up() { + // 1500 rounds to 2000 -> "2k" + assert_eq!(format_context_window(1500), "2k"); + } + + #[test] + fn format_context_window_rounds_down() { + // 1499 rounds to 1000 -> "1k" + assert_eq!(format_context_window(1499), "1k"); + } + + #[test] + fn format_context_window_zero() { + assert_eq!(format_context_window(0), "0"); + } + + // --- format_cost --- + + #[test] + fn format_cost_none() { + assert_eq!(format_cost(None), "-"); + } + + #[test] + fn format_cost_some() { + assert_eq!(format_cost(Some(3.0)), "$3.0"); + } + + #[test] + fn format_cost_fractional() { + assert_eq!(format_cost(Some(15.75)), "$15.8"); + } + + // --- format_speed --- + + #[test] + fn format_speed_none() { + assert_eq!(format_speed(None), "-"); + } + + #[test] + fn format_speed_some() { + assert_eq!(format_speed(Some(85.5)), "85 tok/s"); + } + + // --- resolve_prompt --- + + #[test] + fn resolve_prompt_arg_only() { + let result = resolve_prompt(Some("hello".into()), None).unwrap(); + assert_eq!(result, "hello"); + } + + #[test] + fn resolve_prompt_stdin_only() { + let result = resolve_prompt(None, Some("piped".into())).unwrap(); + assert_eq!(result, "piped"); + } + + #[test] + fn resolve_prompt_both_concatenates() { + let result = resolve_prompt(Some("arg".into()), Some("stdin".into())).unwrap(); + assert_eq!(result, "stdin\narg"); + } + + #[test] + fn resolve_prompt_neither_errors() { + assert!(resolve_prompt(None, None).is_err()); + } + + // --- resolve_model --- + + #[test] + fn resolve_model_explicit_known() { + let (model, provider) = resolve_model(Some("claude-sonnet-4-5".into())); + assert_eq!(model, "claude-sonnet-4-5"); + assert_eq!(provider, Some("anthropic".to_string())); + } + + #[test] + fn resolve_model_explicit_unknown() { + let (model, provider) = resolve_model(Some("nonexistent-model-xyz".into())); + assert_eq!(model, "nonexistent-model-xyz"); + assert_eq!(provider, None); + } + + #[test] + fn resolve_model_none_uses_default() { + let (model, provider) = resolve_model(None); + // Should return some valid model from catalog + assert!(!model.is_empty()); + assert!(provider.is_some()); + } + + // --- apply_options --- + + #[test] + fn apply_options_temperature() { + let params = GenerateParams::new("test-model"); + let result = apply_options(params, &[("temperature".into(), "0.7".into())]).unwrap(); + assert_eq!(result.temperature, Some(0.7)); + } + + #[test] + fn apply_options_max_tokens() { + let params = GenerateParams::new("test-model"); + let result = apply_options(params, &[("max_tokens".into(), "4096".into())]).unwrap(); + assert_eq!(result.max_tokens, Some(4096)); + } + + #[test] + fn apply_options_top_p() { + let params = GenerateParams::new("test-model"); + let result = apply_options(params, &[("top_p".into(), "0.9".into())]).unwrap(); + assert_eq!(result.top_p, Some(0.9)); + } + + #[test] + fn apply_options_unknown_key_goes_to_provider_opts() { + let params = GenerateParams::new("test-model"); + let result = + apply_options(params, &[("custom_key".into(), "custom_val".into())]).unwrap(); + let opts = result.provider_options.unwrap(); + assert_eq!(opts["custom_key"], "custom_val"); + } + + #[test] + fn apply_options_invalid_temperature_errors() { + let params = GenerateParams::new("test-model"); + assert!( + apply_options(params, &[("temperature".into(), "not_a_number".into())]).is_err() + ); + } + + #[test] + fn apply_options_invalid_max_tokens_errors() { + let params = GenerateParams::new("test-model"); + assert!(apply_options(params, &[("max_tokens".into(), "abc".into())]).is_err()); + } + + #[test] + fn apply_options_empty() { + let params = GenerateParams::new("test-model"); + let result = apply_options(params, &[]).unwrap(); + assert_eq!(result.temperature, None); + assert_eq!(result.max_tokens, None); + assert_eq!(result.provider_options, None); + } +} diff --git a/crates/arc-llm/src/provider.rs b/crates/arc-llm/src/provider.rs index 359ecaff5..d4d2e86f9 100644 --- a/crates/arc-llm/src/provider.rs +++ b/crates/arc-llm/src/provider.rs @@ -285,4 +285,72 @@ mod tests { .iter() .all(|p| !p.api_key_env_vars().is_empty())); } + + // Mock adapter that supports all tool choices + struct MockAdapter; + + #[async_trait::async_trait] + impl ProviderAdapter for MockAdapter { + fn name(&self) -> &str { + "mock" + } + async fn complete(&self, _request: &Request) -> Result { + unimplemented!() + } + async fn stream(&self, _request: &Request) -> Result { + unimplemented!() + } + } + + // Mock adapter that rejects "named" tool choice + struct RestrictedAdapter; + + #[async_trait::async_trait] + impl ProviderAdapter for RestrictedAdapter { + fn name(&self) -> &str { + "restricted" + } + async fn complete(&self, _request: &Request) -> Result { + unimplemented!() + } + async fn stream(&self, _request: &Request) -> Result { + unimplemented!() + } + fn supports_tool_choice(&self, mode: &str) -> bool { + mode != "named" + } + } + + #[test] + fn validate_tool_choice_auto_accepted() { + assert!(validate_tool_choice(&MockAdapter, &ToolChoice::Auto).is_ok()); + } + + #[test] + fn validate_tool_choice_none_accepted() { + assert!(validate_tool_choice(&MockAdapter, &ToolChoice::None).is_ok()); + } + + #[test] + fn validate_tool_choice_required_accepted() { + assert!(validate_tool_choice(&MockAdapter, &ToolChoice::Required).is_ok()); + } + + #[test] + fn validate_tool_choice_named_rejected_by_restricted() { + let result = validate_tool_choice(&RestrictedAdapter, &ToolChoice::named("my_tool")); + assert!(result.is_err()); + match result.unwrap_err() { + SdkError::UnsupportedToolChoice { message } => { + assert!(message.contains("restricted")); + assert!(message.contains("named")); + } + other => panic!("expected UnsupportedToolChoice, got {other:?}"), + } + } + + #[test] + fn validate_tool_choice_named_accepted_by_default() { + assert!(validate_tool_choice(&MockAdapter, &ToolChoice::named("my_tool")).is_ok()); + } } diff --git a/crates/arc-llm/src/providers/common.rs b/crates/arc-llm/src/providers/common.rs index 073797704..3a299b417 100644 --- a/crates/arc-llm/src/providers/common.rs +++ b/crates/arc-llm/src/providers/common.rs @@ -271,6 +271,7 @@ impl LineReader { #[cfg(test)] mod tests { use super::*; + use crate::types::ContentPart; #[test] fn is_file_path_absolute() { @@ -379,4 +380,131 @@ mod tests { assert_eq!(info.requests_remaining, None); assert_eq!(info.tokens_limit, Some(10000)); } + + // --- parse_error_body --- + + #[test] + fn parse_error_body_valid_json() { + let body = r#"{"error":{"message":"rate limited","type":"rate_limit_error"}}"#; + let (msg, code, raw) = parse_error_body(body, "type"); + assert_eq!(msg, "rate limited"); + assert_eq!(code.as_deref(), Some("rate_limit_error")); + assert!(raw.is_some()); + } + + #[test] + fn parse_error_body_missing_error_field() { + let body = r#"{"status":"fail"}"#; + let (msg, code, raw) = parse_error_body(body, "type"); + assert_eq!(msg, "Unknown error"); + assert_eq!(code, None); + assert!(raw.is_some()); + } + + #[test] + fn parse_error_body_not_json() { + let body = "Internal Server Error"; + let (msg, code, raw) = parse_error_body(body, "type"); + assert_eq!(msg, "Internal Server Error"); + assert_eq!(code, None); + assert!(raw.is_none()); + } + + #[test] + fn parse_error_body_different_code_field() { + let body = r#"{"error":{"message":"bad","status":"INVALID_ARGUMENT"}}"#; + let (msg, code, _) = parse_error_body(body, "status"); + assert_eq!(msg, "bad"); + assert_eq!(code.as_deref(), Some("INVALID_ARGUMENT")); + } + + #[test] + fn parse_error_body_no_message() { + let body = r#"{"error":{"type":"server_error"}}"#; + let (msg, code, _) = parse_error_body(body, "type"); + assert_eq!(msg, "Unknown error"); + assert_eq!(code.as_deref(), Some("server_error")); + } + + // --- extract_system_prompt --- + + #[test] + fn extract_system_prompt_no_system() { + let msgs = vec![Message::user("hello")]; + let (sys, other) = extract_system_prompt(&msgs); + assert_eq!(sys, None); + assert_eq!(other.len(), 1); + } + + #[test] + fn extract_system_prompt_system_only() { + let msgs = vec![Message::system("Be helpful"), Message::user("hi")]; + let (sys, other) = extract_system_prompt(&msgs); + assert_eq!(sys.as_deref(), Some("Be helpful")); + assert_eq!(other.len(), 1); + assert_eq!(other[0].role, Role::User); + } + + #[test] + fn extract_system_prompt_multiple_system() { + let msgs = vec![ + Message::system("Rule 1"), + Message::system("Rule 2"), + Message::user("hi"), + ]; + let (sys, other) = extract_system_prompt(&msgs); + assert_eq!(sys.as_deref(), Some("Rule 1\nRule 2")); + assert_eq!(other.len(), 1); + } + + #[test] + fn extract_system_prompt_developer_role() { + let dev = Message { + role: Role::Developer, + content: vec![ContentPart::text("dev instructions")], + name: None, + tool_call_id: None, + }; + let msgs = vec![dev, Message::user("hi")]; + let (sys, other) = extract_system_prompt(&msgs); + assert_eq!(sys.as_deref(), Some("dev instructions")); + assert_eq!(other.len(), 1); + } + + #[test] + fn extract_system_prompt_empty() { + let msgs: Vec = vec![]; + let (sys, other) = extract_system_prompt(&msgs); + assert_eq!(sys, None); + assert!(other.is_empty()); + } + + // --- parse_retry_after --- + + #[test] + fn parse_retry_after_valid() { + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert("retry-after", "2.5".parse().unwrap()); + assert_eq!(parse_retry_after(&headers), Some(2.5)); + } + + #[test] + fn parse_retry_after_missing() { + let headers = reqwest::header::HeaderMap::new(); + assert_eq!(parse_retry_after(&headers), None); + } + + #[test] + fn parse_retry_after_invalid() { + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert("retry-after", "not-a-number".parse().unwrap()); + assert_eq!(parse_retry_after(&headers), None); + } + + #[test] + fn parse_retry_after_integer() { + let mut headers = reqwest::header::HeaderMap::new(); + headers.insert("retry-after", "5".parse().unwrap()); + assert_eq!(parse_retry_after(&headers), Some(5.0)); + } } diff --git a/crates/arc-llm/src/tools.rs b/crates/arc-llm/src/tools.rs index bb53227a8..2f19c0128 100644 --- a/crates/arc-llm/src/tools.rs +++ b/crates/arc-llm/src/tools.rs @@ -466,4 +466,36 @@ mod tests { .unwrap() .contains("repair failed")); } + + // --- args_type_name --- + + #[test] + fn args_type_name_null() { + assert_eq!(args_type_name(&serde_json::Value::Null), "null"); + } + + #[test] + fn args_type_name_bool() { + assert_eq!(args_type_name(&serde_json::json!(true)), "boolean"); + } + + #[test] + fn args_type_name_number() { + assert_eq!(args_type_name(&serde_json::json!(42)), "number"); + } + + #[test] + fn args_type_name_string() { + assert_eq!(args_type_name(&serde_json::json!("hello")), "string"); + } + + #[test] + fn args_type_name_array() { + assert_eq!(args_type_name(&serde_json::json!([1, 2])), "array"); + } + + #[test] + fn args_type_name_object() { + assert_eq!(args_type_name(&serde_json::json!({})), "object"); + } } diff --git a/crates/arc-llm/src/types.rs b/crates/arc-llm/src/types.rs index 076becbcb..6f011dc76 100644 --- a/crates/arc-llm/src/types.rs +++ b/crates/arc-llm/src/types.rs @@ -1263,4 +1263,24 @@ mod tests { other => panic!("Expected StepFinish, got {other:?}"), } } + + #[test] + fn tool_choice_mode_str_auto() { + assert_eq!(ToolChoice::Auto.mode_str(), "auto"); + } + + #[test] + fn tool_choice_mode_str_none() { + assert_eq!(ToolChoice::None.mode_str(), "none"); + } + + #[test] + fn tool_choice_mode_str_required() { + assert_eq!(ToolChoice::Required.mode_str(), "required"); + } + + #[test] + fn tool_choice_mode_str_named() { + assert_eq!(ToolChoice::named("get_weather").mode_str(), "named"); + } }