Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
250 changes: 192 additions & 58 deletions crates/goose/src/providers/databricks.rs
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@ struct DatabricksEndpointInfo {
upstream_model_name: Option<String>,
upstream_model_provider: Option<String>,
reasoning: Option<bool>,
supports_responses_api: bool,
}

#[derive(Debug, Clone)]
Expand Down Expand Up @@ -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<DatabricksUpstreamModel> {
Expand Down Expand Up @@ -304,6 +313,7 @@ impl DatabricksProvider {

fn endpoint_info_from_value(endpoint: &Value) -> Option<DatabricksEndpointInfo> {
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);
Expand All @@ -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,
Expand Down Expand Up @@ -378,13 +424,16 @@ impl DatabricksProvider {
let mut current_endpoint_name = endpoint_name.to_string();
let mut visited = HashSet::new();
let mut last_info: Option<DatabricksEndpointInfo> = None;
let mut first_hop_supports_responses_api: Option<bool> = None;

for _ in 0..MAX_MODEL_SERVING_HOPS {
if !visited.insert(current_endpoint_name.clone()) {
break;
}

let info = self.fetch_endpoint_info(&current_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(),
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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(
Expand Down Expand Up @@ -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)
}
}

Expand All @@ -514,7 +570,10 @@ impl DatabricksProvider {
) -> Result<Value, ProviderError> {
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();

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Skip metadata lookup for embedding posts

When create_embeddings calls post, the payload is an embedding request (input without messages), and get_endpoint_path(..., is_embedding, ...) ignores the responses flag and always returns the hard-coded embedding endpoint. This unconditional metadata fetch therefore makes the first embedding call block on GET /api/2.0/serving-endpoints/{self.model.model_name} even though the result is unused; if the configured chat endpoint's metadata is slow or unavailable while the embedding endpoint is healthy, embeddings can stall until the provider HTTP timeout before the actual embedding request is sent.

Useful? React with 👍 / 👎.

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) {
Expand Down Expand Up @@ -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);

Expand Down Expand Up @@ -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!({
Expand All @@ -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");
Expand Down Expand Up @@ -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"]
));
}
}
Loading