From b9b671c6c3f6a019530698773af2354afece3756 Mon Sep 17 00:00:00 2001 From: Kalvin C Date: Wed, 26 Aug 2026 15:48:55 +0000 Subject: [PATCH] feat(provider): discover Databricks Unity Catalog model services (#11575) --- crates/goose-providers/src/databricks_v2.rs | 552 ++++++++++++++++++-- crates/goose/src/providers/inventory/mod.rs | 41 +- 2 files changed, 534 insertions(+), 59 deletions(-) diff --git a/crates/goose-providers/src/databricks_v2.rs b/crates/goose-providers/src/databricks_v2.rs index 055767652..0a2e5ccd0 100644 --- a/crates/goose-providers/src/databricks_v2.rs +++ b/crates/goose-providers/src/databricks_v2.rs @@ -1,5 +1,6 @@ use crate::formats::anthropic::{AnthropicFormatOptions, ANTHROPIC_PROVIDER_NAME}; use crate::formats::openai::{self, extract_reasoning_effort, is_openai_responses_model}; +use crate::http_status::{read_error_body, read_json_response}; use crate::images::ImageFormat; use anyhow::Result; use async_stream::try_stream; @@ -7,6 +8,7 @@ use async_trait::async_trait; use futures::TryStreamExt; use serde::Serialize; use serde_json::Value; +use std::collections::HashSet; use std::io; use std::sync::{Arc, Mutex}; use std::time::Duration; @@ -36,7 +38,40 @@ use rmcp::model::Tool; const DATABRICKS_V2_PROVIDER_NAME: &str = "databricks_v2"; const DATABRICKS_V2_LIST_ENDPOINTS_PATH: &str = "api/ai-gateway/v2/endpoints"; -const DATABRICKS_V2_LIST_ENDPOINTS_PAGE_SIZE: usize = 100; +const DATABRICKS_V2_LIST_MODEL_SERVICES_PATH: &str = "api/2.1/unity-catalog/model-services"; +const DATABRICKS_V2_MODEL_SERVICE_PREFIX: &str = "model-services/"; +const DATABRICKS_V2_CATALOG_PAGE_SIZE: usize = 100; +const DATABRICKS_V2_MAX_CATALOG_PAGES: usize = 100; +// Model-services intermittently uses 499 for transient gateway timeouts. +const DATABRICKS_V2_TRANSIENT_GATEWAY_STATUS: u16 = 499; + +#[derive(Clone, Copy)] +struct ModelCatalog { + path: &'static str, + items_key: &'static str, + name_prefix: Option<&'static str>, + view: Option<&'static str>, + label: &'static str, +} + +const DATABRICKS_V2_ENDPOINTS_CATALOG: ModelCatalog = ModelCatalog { + path: DATABRICKS_V2_LIST_ENDPOINTS_PATH, + items_key: "endpoints", + name_prefix: None, + view: None, + label: "AI Gateway endpoints", +}; + +// LIST can include metadata-only services but omits caller-effective grants. +// Inference enforces EXECUTE. +const DATABRICKS_V2_MODEL_SERVICES_CATALOG: ModelCatalog = ModelCatalog { + path: DATABRICKS_V2_LIST_MODEL_SERVICES_PATH, + items_key: "model_services", + name_prefix: Some(DATABRICKS_V2_MODEL_SERVICE_PREFIX), + view: Some("FULL"), + label: "model services", +}; + pub const DATABRICKS_V2_DEFAULT_MODEL: &str = "databricks-gpt-5-5"; pub const DATABRICKS_V2_KNOWN_MODELS: &[&str] = &["databricks-gpt-5-5", "databricks-claude-opus-4-7"]; @@ -137,6 +172,10 @@ impl DatabricksV2Provider { } fn route_for_model(model_name: &str) -> DatabricksV2Route { + if Self::is_model_service_fqn(model_name) { + // UC namespaces are user-defined and cannot select a native API. + return DatabricksV2Route::MlflowChatCompletions; + } let (clean_name, _) = extract_reasoning_effort(model_name); let lower = clean_name.to_lowercase(); @@ -149,6 +188,16 @@ impl DatabricksV2Provider { } } + fn is_model_service_fqn(model_name: &str) -> bool { + let Some((catalog, remainder)) = model_name.split_once('.') else { + return false; + }; + let Some((schema, service)) = remainder.split_once('.') else { + return false; + }; + !catalog.is_empty() && !schema.is_empty() && !service.is_empty() + } + fn looks_like_gpt5(model_name: &str) -> bool { model_name.contains("gpt-5") || model_name.contains("gpt5") } @@ -157,26 +206,59 @@ impl DatabricksV2Provider { model_name.contains("claude") } - fn parse_list_endpoints_response( + fn name_looks_chat_capable(name: &str) -> bool { + if name.to_ascii_lowercase().contains("embedding") { + return false; + } + !name + .split(|c: char| !c.is_ascii_alphanumeric()) + .any(|segment| { + segment.eq_ignore_ascii_case("bge") || segment.eq_ignore_ascii_case("gte") + }) + } + + fn model_service_supports_chat(item: &Value, fallback_name: &str) -> bool { + let Some(api_types) = item.get("supported_api_types") else { + // Older workspaces omit capabilities; only the service leaf is safe to inspect. + return Self::name_looks_chat_capable(fallback_name); + }; + let Some(api_types) = api_types.as_array() else { + return false; + }; + + api_types.iter().filter_map(Value::as_str).any(|api_type| { + api_type.eq_ignore_ascii_case("chat") + || api_type.eq_ignore_ascii_case("mlflow/v1/chat/completions") + }) + } + + fn parse_catalog_page( json: &Value, + catalog: &ModelCatalog, ) -> Result<(Vec, Option), ProviderError> { - let endpoints = json - .get("endpoints") + let items = json + .get(catalog.items_key) .and_then(|v| v.as_array()) .ok_or_else(|| { - ProviderError::RequestFailed( - "Unexpected response format from Databricks AI Gateway endpoints API" - .to_string(), - ) + ProviderError::RequestFailed(format!( + "Unexpected response format from Databricks {} API", + catalog.label + )) })?; - let models: Vec = endpoints + let models: Vec = items .iter() - .filter_map(|endpoint| { - endpoint - .get("name") - .and_then(|v| v.as_str()) - .map(str::to_string) + .filter_map(|item| { + let name = item.get("name").and_then(Value::as_str)?; + let (name, is_chat_capable) = match catalog.name_prefix { + Some(prefix) => { + let name = name.strip_prefix(prefix)?; + let leaf = name.rsplit('.').next().unwrap_or(name); + (name, Self::model_service_supports_chat(item, leaf)) + } + None => (name, Self::name_looks_chat_capable(name)), + }; + (!name.is_empty() && is_chat_capable).then(|| name.to_string()) }) .collect(); @@ -227,14 +309,23 @@ impl DatabricksV2Provider { messages: &[Message], tools: &[Tool], ) -> Result { + let is_model_service = Self::is_model_service_fqn(&model_config.model_name); + let mut format_config = model_config.clone(); + if is_model_service { + // Keep UC namespace text out of OpenAI format heuristics. + format_config.model_name = "model-service".to_string(); + } let mut payload = openai::create_request( - model_config, + &format_config, system, messages, tools, &ImageFormat::OpenAi, true, )?; + if is_model_service { + payload["model"] = Value::String(model_config.model_name.clone()); + } if payload.get("max_tokens").is_none() { payload["max_tokens"] = Value::from(model_config.max_output_tokens()); } @@ -370,49 +461,105 @@ impl Provider for DatabricksV2Provider { } async fn fetch_supported_models(&self) -> Result, ProviderError> { - let mut models = Vec::new(); - let mut page_token: Option = None; + let (endpoint_result, service_result) = tokio::join!( + self.fetch_model_catalog(&DATABRICKS_V2_ENDPOINTS_CATALOG), + self.fetch_model_catalog(&DATABRICKS_V2_MODEL_SERVICES_CATALOG), + ); - loop { - let mut path = format!( - "{}?page_size={}", - DATABRICKS_V2_LIST_ENDPOINTS_PATH, DATABRICKS_V2_LIST_ENDPOINTS_PAGE_SIZE - ); - if let Some(token) = &page_token { - path.push_str(&format!("&page_token={}", urlencoding::encode(token))); + let mut names = Vec::new(); + let mut failures = Vec::new(); + let mut any_catalog_succeeded = false; + for (catalog, result) in [ + (&DATABRICKS_V2_ENDPOINTS_CATALOG, endpoint_result), + (&DATABRICKS_V2_MODEL_SERVICES_CATALOG, service_result), + ] { + match result { + Ok(models) => { + any_catalog_succeeded = true; + names.extend(models); + } + Err(error) => failures.push((catalog.label, error)), } - - let response = self.api_client.response_get(&path).await.map_err(|e| { - ProviderError::RequestFailed(format!( - "Failed to fetch Databricks AI Gateway endpoints: {e}" - )) - })?; - - if !response.status().is_success() { - let status = response.status(); - let detail = response.text().await.unwrap_or_default(); - return Err(ProviderError::RequestFailed(format!( - "Failed to fetch Databricks AI Gateway endpoints: {status} {detail}" - ))); - } - - let json: Value = response.json().await.map_err(|e| { - ProviderError::RequestFailed(format!( - "Failed to parse Databricks AI Gateway endpoints response: {e}" - )) - })?; - - let (page_models, next_page_token) = Self::parse_list_endpoints_response(&json)?; - models.extend(page_models); - - if next_page_token.is_none() || next_page_token == page_token { - break; - } - page_token = next_page_token; } - models.sort(); - Ok(models) + if !any_catalog_succeeded { + let details = failures + .into_iter() + .map(|(_, error)| match error { + ProviderError::RequestFailed(message) => message, + error => error.to_string(), + }) + .collect::>() + .join("; "); + return Err(ProviderError::RequestFailed(details)); + } + for (label, error) in failures { + tracing::warn!(catalog = label, %error, "Failed to fetch Databricks model catalog"); + } + + names.sort(); + names.dedup(); + Ok(names) + } +} + +impl DatabricksV2Provider { + async fn fetch_model_catalog( + &self, + catalog: &ModelCatalog, + ) -> Result, ProviderError> { + let ModelCatalog { path, label, .. } = catalog; + let mut models = Vec::new(); + let mut page_token: Option = None; + let mut seen_page_tokens = HashSet::new(); + + for _ in 0..DATABRICKS_V2_MAX_CATALOG_PAGES { + let mut path_with_query = format!("{path}?page_size={DATABRICKS_V2_CATALOG_PAGE_SIZE}"); + if let Some(view) = catalog.view { + path_with_query.push_str(&format!("&view={}", urlencoding::encode(view))); + } + if let Some(token) = &page_token { + path_with_query.push_str(&format!("&page_token={}", urlencoding::encode(token))); + } + + let json: Value = self + .with_retry_config( + || async { + let response = self.api_client.response_get(&path_with_query).await?; + if response.status().as_u16() == DATABRICKS_V2_TRANSIENT_GATEWAY_STATUS { + let detail = read_error_body(response).await.unwrap_or_default(); + return Err(ProviderError::ServerError(format!( + "Databricks {label} returned {DATABRICKS_V2_TRANSIENT_GATEWAY_STATUS}: {detail}" + ))); + } + read_json_response(handle_status(response).await?).await + }, + self.retry_config.clone().transient_only(), + ) + .await + .map_err(|error| { + ProviderError::RequestFailed(format!( + "Failed to fetch Databricks {label}: {error}" + )) + })?; + + let (page_models, next_page_token) = Self::parse_catalog_page(&json, catalog)?; + models.extend(page_models); + + let Some(next_page_token) = next_page_token else { + return Ok(models); + }; + if !seen_page_tokens.insert(next_page_token.clone()) { + return Err(ProviderError::RequestFailed(format!( + "Databricks {label} returned a repeated page token" + ))); + } + page_token = Some(next_page_token); + } + + Err(ProviderError::RequestFailed(format!( + "Databricks {label} pagination exceeded {DATABRICKS_V2_MAX_CATALOG_PAGES} pages" + ))) } } @@ -456,7 +603,8 @@ mod tests { }); let (models, next_page_token) = - DatabricksV2Provider::parse_list_endpoints_response(&json).unwrap(); + DatabricksV2Provider::parse_catalog_page(&json, &DATABRICKS_V2_ENDPOINTS_CATALOG) + .unwrap(); assert_eq!( models, @@ -473,11 +621,303 @@ mod tests { fn errors_when_list_endpoints_response_has_no_endpoints_array() { let json = serde_json::json!({"data": []}); - let error = DatabricksV2Provider::parse_list_endpoints_response(&json).unwrap_err(); + let error = + DatabricksV2Provider::parse_catalog_page(&json, &DATABRICKS_V2_ENDPOINTS_CATALOG) + .unwrap_err(); assert!(matches!(error, ProviderError::RequestFailed(_))); assert!(error .to_string() .contains("Unexpected response format from Databricks AI Gateway endpoints API")); } + + #[test] + fn filters_non_chat_endpoints() { + let json = serde_json::json!({ + "endpoints": [ + {"name": "databricks-bge-large-en"}, + {"name": "my-gte-small"}, + {"name": "text-embedding-3-large"}, + {"name": "databricks-gpt-5-5"}, + {"name": "custom-model"} + ] + }); + + let (models, _) = + DatabricksV2Provider::parse_catalog_page(&json, &DATABRICKS_V2_ENDPOINTS_CATALOG) + .unwrap(); + + assert_eq!( + models, + vec!["databricks-gpt-5-5".to_string(), "custom-model".to_string(),] + ); + } + + #[test] + fn parses_and_filters_model_services() { + let json = serde_json::json!({ + "model_services": [ + { + "name": "model-services/catalog.schema.vector-search", + "supported_api_types": ["mlflow/v1/embeddings"] + }, + { + "name": "model-services/catalog.schema.embedding-assistant", + "supported_api_types": ["mlflow/v1/chat/completions"] + }, + {"name": "model-services/embedding_catalog.gte_schema.chat-model"}, + {"name": "model-services/catalog.schema.bge-embedding"}, + {"name": "model-services/"}, + {"name": "no-prefix"} + ] + }); + + let (models, _) = + DatabricksV2Provider::parse_catalog_page(&json, &DATABRICKS_V2_MODEL_SERVICES_CATALOG) + .unwrap(); + + assert_eq!( + models, + vec![ + "catalog.schema.embedding-assistant", + "embedding_catalog.gte_schema.chat-model", + ] + ); + } + + #[test] + fn routes_model_service_fqns_to_mlflow() { + for model in ["team.claude.kimi-chat", "gpt-5.schema.kimi-chat"] { + assert_eq!( + DatabricksV2Provider::route_for_model(model), + DatabricksV2Route::MlflowChatCompletions, + "unexpected route for {model}" + ); + } + } + + mod fetch_supported_models { + use super::*; + use serde_json::json; + use wiremock::matchers::{method, path, query_param, query_param_is_missing}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + fn provider(server: &MockServer) -> DatabricksV2Provider { + provider_with_retry(server, RetryConfig::new(0, 0, 1.0, 0)) + } + + fn provider_with_retry( + server: &MockServer, + retry_config: RetryConfig, + ) -> DatabricksV2Provider { + DatabricksV2Provider::new( + server.uri(), + DatabricksAuth::token("test-token".to_string()), + retry_config, + None, + None, + None, + None, + None, + ) + .unwrap() + } + + async fn mount_endpoints(server: &MockServer, body: serde_json::Value) { + Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_ENDPOINTS_PATH}"))) + .and(query_param("page_size", "100")) + .respond_with(ResponseTemplate::new(200).set_body_json(body)) + .expect(1) + .mount(server) + .await; + } + + #[tokio::test] + async fn returns_union_of_both_catalogs_sorted_and_deduplicated() { + let server = MockServer::start().await; + mount_endpoints( + &server, + json!({"endpoints": [ + {"name": "databricks-gpt-5-5"}, + {"name": "catalog.schema.shared-model"} + ]}), + ) + .await; + Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_MODEL_SERVICES_PATH}"))) + .and(query_param("page_size", "100")) + .and(query_param("view", "FULL")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "model_services": [ + {"name": "model-services/catalog.schema.shared-model"}, + {"name": "model-services/data.goose.goose-kimi-k3"} + ] + }))) + .expect(1) + .mount(&server) + .await; + + let models = provider(&server).fetch_supported_models().await.unwrap(); + + assert_eq!( + models, + vec![ + "catalog.schema.shared-model", + "data.goose.goose-kimi-k3", + "databricks-gpt-5-5", + ] + ); + } + + #[tokio::test] + async fn paginates_model_services_with_url_encoded_tokens() { + let server = MockServer::start().await; + mount_endpoints(&server, json!({"endpoints": [{"name": "endpoint"}]})).await; + + for (page_token, name, next_page_token) in [ + (None, "a.b.c", Some("svc tok%")), + (Some("svc tok%"), "a.b.d", Some("token-b")), + (Some("token-b"), "a.b.e", None), + ] { + let mut mock = Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_MODEL_SERVICES_PATH}"))) + .and(query_param("page_size", "100")) + .and(query_param("view", "FULL")); + mock = match page_token { + Some(page_token) => mock.and(query_param("page_token", page_token)), + None => mock.and(query_param_is_missing("page_token")), + }; + mock.respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "model_services": [{"name": format!("model-services/{name}")}], + "next_page_token": next_page_token + }))) + .expect(1) + .mount(&server) + .await; + } + + let models = provider(&server).fetch_supported_models().await.unwrap(); + + assert_eq!(models, vec!["a.b.c", "a.b.d", "a.b.e", "endpoint"]); + } + + #[tokio::test] + async fn rejects_model_service_page_token_cycles() { + let server = MockServer::start().await; + for page_token in [None, Some("token-a")] { + let mut mock = Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_MODEL_SERVICES_PATH}"))); + mock = match page_token { + Some(page_token) => mock.and(query_param("page_token", page_token)), + None => mock.and(query_param_is_missing("page_token")), + }; + mock.respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "model_services": [], + "next_page_token": "token-a" + }))) + .expect(1) + .mount(&server) + .await; + } + + let error = provider(&server) + .fetch_model_catalog(&DATABRICKS_V2_MODEL_SERVICES_CATALOG) + .await + .unwrap_err(); + + assert!(error.to_string().contains("repeated page token")); + } + + #[tokio::test] + async fn does_not_retry_permanent_catalog_failures() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_MODEL_SERVICES_PATH}"))) + .respond_with(ResponseTemplate::new(404).set_body_json(json!({ + "message": "not available" + }))) + .expect(1) + .mount(&server) + .await; + + let error = provider_with_retry(&server, RetryConfig::new(3, 0, 1.0, 0)) + .fetch_model_catalog(&DATABRICKS_V2_MODEL_SERVICES_CATALOG) + .await + .unwrap_err(); + + assert!(error.to_string().contains("not available")); + } + + #[tokio::test] + async fn retries_transient_catalog_failures() { + for status in [500, DATABRICKS_V2_TRANSIENT_GATEWAY_STATUS] { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_MODEL_SERVICES_PATH}"))) + .respond_with(ResponseTemplate::new(status)) + .up_to_n_times(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_MODEL_SERVICES_PATH}"))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "model_services": [] + }))) + .expect(1) + .mount(&server) + .await; + + let models = provider_with_retry(&server, RetryConfig::new(1, 0, 1.0, 0)) + .fetch_model_catalog(&DATABRICKS_V2_MODEL_SERVICES_CATALOG) + .await + .unwrap(); + + assert!(models.is_empty()); + } + } + + #[tokio::test] + async fn returns_endpoints_when_services_fail() { + let server = MockServer::start().await; + mount_endpoints(&server, json!({"endpoints": [{"name": "only-endpoint"}]})).await; + Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_MODEL_SERVICES_PATH}"))) + .respond_with(ResponseTemplate::new(500).set_body_json(json!({"message": "boom"}))) + .mount(&server) + .await; + + let models = provider(&server).fetch_supported_models().await.unwrap(); + + assert_eq!(models, vec!["only-endpoint"]); + } + + #[tokio::test] + async fn errors_when_both_catalogs_fail() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_ENDPOINTS_PATH}"))) + .respond_with( + ResponseTemplate::new(500).set_body_json(json!({"message": "endpoints down"})), + ) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path(format!("/{DATABRICKS_V2_LIST_MODEL_SERVICES_PATH}"))) + .respond_with( + ResponseTemplate::new(500).set_body_json(json!({"message": "services down"})), + ) + .mount(&server) + .await; + + let err = provider(&server) + .fetch_supported_models() + .await + .unwrap_err(); + + assert!(matches!(err, ProviderError::RequestFailed(_))); + assert!(err.to_string().contains("endpoints down")); + assert!(err.to_string().contains("services down")); + } + } } diff --git a/crates/goose/src/providers/inventory/mod.rs b/crates/goose/src/providers/inventory/mod.rs index 66b9b6e87..4dfb3648b 100644 --- a/crates/goose/src/providers/inventory/mod.rs +++ b/crates/goose/src/providers/inventory/mod.rs @@ -1041,6 +1041,10 @@ fn fallback_inventory_identity(provider_id: &str) -> InventoryIdentityInput { ) } +fn is_databricks_v2_model_service(provider_family: &str, model_id: &str) -> bool { + provider_family == "databricks_v2" && model_id.splitn(3, '.').count() == 3 +} + fn enrich_model_ids_with_canonical( provider_family: &str, model_ids: &[String], @@ -1060,11 +1064,16 @@ fn enrich_model_ids_with_canonical( } let mut models: Vec = Vec::new(); - let mut seen_names: HashSet = HashSet::new(); + let mut seen_keys: HashSet = HashSet::new(); for id in model_ids { let model = enriched_model(provider_family, id, None); - if !seen_names.insert(model.name.clone()) { + let dedup_key = if is_databricks_v2_model_service(provider_family, id) { + id + } else { + &model.name + }; + if !seen_keys.insert(dedup_key.clone()) { continue; } models.push(model); @@ -1076,7 +1085,9 @@ fn enrich_model_ids_with_canonical( if matches!(provider_family, "databricks" | "databricks_v2") { let mut name_to_idx: HashMap = HashMap::new(); for (idx, model) in models.iter().enumerate() { - name_to_idx.insert(model.name.clone(), idx); + if !is_databricks_v2_model_service(provider_family, &model.id) { + name_to_idx.insert(model.name.clone(), idx); + } } for id in model_ids { if !id.starts_with("goose-") { @@ -1349,6 +1360,30 @@ mod tests { ); } + #[test] + fn databricks_v2_inventory_preserves_distinct_model_service_fqns() { + let model_ids = [ + "alpha.prod.claude-sonnet-4-5", + "beta.prod.claude-sonnet-4-5", + "alpha.prod.claude-sonnet-4-5", + ] + .map(String::from); + let models = enrich_model_ids_with_canonical("databricks_v2", &model_ids); + + let ids = models + .iter() + .map(|model| model.id.as_str()) + .collect::>(); + assert_eq!( + ids, + [ + "alpha.prod.claude-sonnet-4-5", + "beta.prod.claude-sonnet-4-5" + ] + ); + assert_eq!(models[0].name, models[1].name); + } + #[test] fn inventory_uses_configured_models_before_first_successful_refresh() { let configured_models = [ModelInfo::new("claude-sonnet-4-5").with_context_limit(0)];