diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index 35d8afd8aff1..6e1304237fc5 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -44,6 +44,7 @@ struct DatabricksEndpointInfo { upstream_model_name: Option, upstream_model_provider: Option, reasoning: Option, + supports_responses_api: bool, } #[derive(Debug, Clone)] @@ -234,16 +235,24 @@ impl DatabricksProvider { } } - fn is_responses_model(model_name: &str) -> bool { - is_openai_responses_model(model_name) - } - fn is_claude_model(model_name: &str) -> bool { model_name.to_lowercase().contains("claude") } fn is_reasoning_capable_model_name(model_name: &str) -> bool { - Self::is_claude_model(model_name) || Self::is_responses_model(model_name) + Self::is_claude_model(model_name) || is_openai_responses_model(model_name) + } + + fn uses_responses_api( + endpoint_info: Option<&DatabricksEndpointInfo>, + model_names: &[&str], + ) -> bool { + match endpoint_info { + Some(info) => info.supports_responses_api, + None => model_names + .iter() + .any(|name| is_openai_responses_model(name)), + } } fn endpoint_model_candidates(value: &Value) -> Vec { @@ -304,6 +313,7 @@ impl DatabricksProvider { fn endpoint_info_from_value(endpoint: &Value) -> Option { let name = endpoint.get("name")?.as_str()?.to_string(); + let supports_responses_api = Self::endpoint_supports_responses_api(endpoint); let upstream_model = Self::endpoint_model_candidates(endpoint) .into_iter() .find(|candidate| candidate.name != name); @@ -320,9 +330,45 @@ impl DatabricksProvider { upstream_model_name, upstream_model_provider, reasoning, + supports_responses_api, }) } + fn endpoint_supports_responses_api(endpoint: &Value) -> bool { + fn value_contains_responses_api(value: &Value) -> bool { + match value { + Value::Object(map) => { + map.get("api_types") + .and_then(|api_types| api_types.as_array()) + .is_some_and(|api_types| { + api_types + .iter() + .any(|api_type| api_type.as_str() == Some("openai/v1/responses")) + }) + || map.values().any(value_contains_responses_api) + } + Value::Array(values) => values.iter().any(value_contains_responses_api), + _ => false, + } + } + + let Some(config) = endpoint.get("config") else { + return false; + }; + + for collection_key in ["served_entities", "served_models"] { + let Some(entities) = config.get(collection_key).and_then(|v| v.as_array()) else { + continue; + }; + + if entities.iter().any(value_contains_responses_api) { + return true; + } + } + + false + } + async fn fetch_endpoint_info( &self, endpoint_name: &str, @@ -378,6 +424,7 @@ impl DatabricksProvider { let mut current_endpoint_name = endpoint_name.to_string(); let mut visited = HashSet::new(); let mut last_info: Option = None; + let mut first_hop_supports_responses_api: Option = None; for _ in 0..MAX_MODEL_SERVING_HOPS { if !visited.insert(current_endpoint_name.clone()) { @@ -385,6 +432,8 @@ impl DatabricksProvider { } let info = self.fetch_endpoint_info(¤t_endpoint_name).await?; + let supports_responses_api = + *first_hop_supports_responses_api.get_or_insert(info.supports_responses_api); let next_endpoint_name = match ( info.upstream_model_provider.as_deref(), info.upstream_model_name.as_deref(), @@ -403,7 +452,7 @@ impl DatabricksProvider { continue; } - return Ok(if info.name == original_endpoint_name { + let mut resolved_info = if info.name == original_endpoint_name { info } else { let upstream_model_name = info @@ -415,8 +464,11 @@ impl DatabricksProvider { upstream_model_name, upstream_model_provider: info.upstream_model_provider.clone(), reasoning: info.reasoning, + supports_responses_api, } - }); + }; + resolved_info.supports_responses_api = supports_responses_api; + return Ok(resolved_info); } last_info @@ -425,6 +477,7 @@ impl DatabricksProvider { upstream_model_name: info.upstream_model_name, upstream_model_provider: info.upstream_model_provider, reasoning: info.reasoning, + supports_responses_api: first_hop_supports_responses_api.unwrap_or(false), }) .ok_or_else(|| { ProviderError::RequestFailed( @@ -484,16 +537,19 @@ impl DatabricksProvider { } } - fn get_endpoint_path(&self, model_name: &str, is_embedding: bool) -> String { + fn get_endpoint_path( + &self, + model_name: &str, + is_embedding: bool, + is_responses_model: bool, + ) -> String { if is_embedding { "serving-endpoints/text-embedding-3-small/invocations".to_string() + } else if is_responses_model { + "serving-endpoints/responses".to_string() } else { let (clean_name, _) = extract_reasoning_effort(model_name); - if Self::is_responses_model(&clean_name) { - "serving-endpoints/responses".to_string() - } else { - format!("serving-endpoints/{}/invocations", clean_name) - } + format!("serving-endpoints/{}/invocations", clean_name) } } @@ -514,7 +570,10 @@ impl DatabricksProvider { ) -> Result { let is_embedding = payload.get("input").is_some() && payload.get("messages").is_none(); let model_to_use = model_name.unwrap_or(&self.model.model_name); - let path = self.get_endpoint_path(model_to_use, is_embedding); + let (endpoint_name, _) = extract_reasoning_effort(model_to_use); + let endpoint_info = self.resolve_endpoint_info_cached(&endpoint_name).await.ok(); + let is_responses_model = Self::uses_responses_api(endpoint_info.as_ref(), &[model_to_use]); + let path = self.get_endpoint_path(model_to_use, is_embedding, is_responses_model); if let Some(session_id) = session_id { if let Some(client_request_id) = self.build_client_request_id(session_id) { @@ -595,12 +654,14 @@ impl Provider for DatabricksProvider { .as_ref() .and_then(|info| info.upstream_model_name.as_deref()) .unwrap_or(&model_config.model_name); - let is_responses_model = Self::is_responses_model(&model_config.model_name) - || Self::is_responses_model(effective_model_name); + let is_responses_model = Self::uses_responses_api( + endpoint_info.as_ref(), + &[&model_config.model_name, effective_model_name], + ); let path = if is_responses_model { "serving-endpoints/responses".to_string() } else { - self.get_endpoint_path(&model_config.model_name, false) + self.get_endpoint_path(&model_config.model_name, false, is_responses_model) }; let client_request_id = self.build_client_request_id(session_id); @@ -880,47 +941,6 @@ impl EmbeddingCapable for DatabricksProvider { mod tests { use super::*; - fn test_provider() -> DatabricksProvider { - DatabricksProvider { - api_client: super::super::api_client::ApiClient::new( - "https://example.com".to_string(), - super::super::api_client::AuthMethod::NoAuth, - ) - .unwrap(), - host: "https://example.com".to_string(), - auth: DatabricksAuth::Token("fake".into()), - model: ModelConfig::new_or_fail("databricks-gpt-5.4"), - image_format: ImageFormat::OpenAi, - retry_config: RetryConfig::default(), - fast_retry_config: RetryConfig::new(0, 0, 1.0, 0), - name: "databricks".into(), - token_cache: std::sync::Arc::new(std::sync::Mutex::new(None)), - instance_id: None, - } - } - - #[test] - fn responses_models_route_to_responses_endpoint() { - let provider = test_provider(); - - for (model_name, expected_path) in [ - ("gpt-5.4", "serving-endpoints/responses"), - ("databricks-gpt-5.4-high", "serving-endpoints/responses"), - ("databricks-gpt-5-4-xhigh", "serving-endpoints/responses"), - ("o3-mini", "serving-endpoints/responses"), - ( - "databricks-claude-sonnet-4", - "serving-endpoints/databricks-claude-sonnet-4/invocations", - ), - ] { - assert_eq!( - provider.get_endpoint_path(model_name, false), - expected_path, - "unexpected endpoint for {model_name}" - ); - } - } - #[test] fn endpoint_metadata_marks_reasoning_alias_from_external_model() { let endpoint = json!({ @@ -942,6 +962,7 @@ mod tests { assert_eq!(info.name, "goose"); assert_eq!(info.upstream_model_name.as_deref(), Some("claude-opus-4.6")); assert_eq!(info.reasoning, Some(true)); + assert!(!info.supports_responses_api); let model_info = DatabricksProvider::model_info_from_endpoint(info); assert_eq!(model_info.name, "goose"); @@ -1014,5 +1035,118 @@ mod tests { assert_eq!(info.name, "goose-gpt-5-5"); assert_eq!(info.upstream_model_name, None); assert_eq!(info.reasoning, Some(true)); + assert!(!info.supports_responses_api); + } + + #[test] + fn endpoint_metadata_detects_responses_api_from_foundation_model_api_types() { + let endpoint = json!({ + "name": "databricks-gpt-5-4", + "config": { + "served_entities": [{ + "name": "databricks-gpt-5-4", + "entity_name": "system.ai.databricks-gpt-5-4", + "type": "FOUNDATION_MODEL", + "foundation_model": { + "name": "system.ai.databricks-gpt-5-4", + "display_name": "GPT-5.4", + "api_types": [ + "mlflow/v1/chat/completions", + "openai/v1/responses", + "cursor/v1/chat/completions" + ] + } + }] + } + }); + + let info = DatabricksProvider::endpoint_info_from_value(&endpoint).unwrap(); + + assert_eq!(info.name, "databricks-gpt-5-4"); + assert_eq!( + info.upstream_model_name.as_deref(), + Some("system.ai.databricks-gpt-5-4") + ); + assert!(info.supports_responses_api); + } + + #[test] + fn endpoint_metadata_detects_responses_api_from_served_models() { + let endpoint = json!({ + "name": "databricks-gpt-5-4", + "config": { + "served_models": [{ + "foundation_model": { + "api_types": ["openai/v1/responses"] + } + }] + } + }); + + let info = DatabricksProvider::endpoint_info_from_value(&endpoint).unwrap(); + + assert!(info.supports_responses_api); + } + + #[test] + fn endpoint_metadata_ignores_pending_config_for_responses_routing() { + let endpoint = json!({ + "name": "databricks-gpt-5-4", + "config": { + "served_entities": [{ + "foundation_model": { + "api_types": ["mlflow/v1/chat/completions"] + } + }] + }, + "pending_config": { + "served_entities": [{ + "foundation_model": { + "api_types": ["openai/v1/responses"] + } + }] + } + }); + + let info = DatabricksProvider::endpoint_info_from_value(&endpoint).unwrap(); + + assert!(!info.supports_responses_api); + } + + #[test] + fn responses_routing_prefers_metadata_over_model_name() { + let responses_info = DatabricksEndpointInfo { + name: "custom".into(), + upstream_model_name: None, + upstream_model_provider: None, + reasoning: None, + supports_responses_api: true, + }; + assert!(DatabricksProvider::uses_responses_api( + Some(&responses_info), + &["databricks-claude-sonnet-4"] + )); + + let chat_info = DatabricksEndpointInfo { + supports_responses_api: false, + ..responses_info + }; + assert!(!DatabricksProvider::uses_responses_api( + Some(&chat_info), + &["gpt-5.4"] + )); + } + + #[test] + fn responses_routing_falls_back_to_model_name_without_metadata() { + assert!(DatabricksProvider::uses_responses_api(None, &["gpt-5.4"])); + assert!(DatabricksProvider::uses_responses_api( + None, + &["databricks-claude-sonnet-4", "gpt-5.4"] + )); + assert!(!DatabricksProvider::uses_responses_api( + None, + &["databricks-claude-sonnet-4"] + )); } }