fix: Parse maas models for gcp vertex provider (#5816)

This commit is contained in:
David Katz
2025-11-20 10:51:36 -05:00
committed by GitHub
parent cfdf01567d
commit c9e9af52c6
2 changed files with 39 additions and 7 deletions
@@ -68,6 +68,9 @@ pub enum GcpVertexAIModel {
Claude(ClaudeVersion), Claude(ClaudeVersion),
/// Gemini model family with specific versions /// Gemini model family with specific versions
Gemini(GeminiVersion), Gemini(GeminiVersion),
/// MaaS (Model as a Service) models from Model Garden
/// Contains (publisher, full_model_name)
MaaS(String, String),
} }
/// Represents available versions of the Claude model for goose. /// Represents available versions of the Claude model for goose.
@@ -126,6 +129,7 @@ impl fmt::Display for GcpVertexAIModel {
GeminiVersion::Pro25 => "gemini-2.5-pro", GeminiVersion::Pro25 => "gemini-2.5-pro",
GeminiVersion::Generic(name) => name, GeminiVersion::Generic(name) => name,
}, },
Self::MaaS(_, model_name) => model_name,
}; };
write!(f, "{model_id}") write!(f, "{model_id}")
} }
@@ -137,10 +141,12 @@ impl GcpVertexAIModel {
/// Each model family has a well-known location based on availability: /// Each model family has a well-known location based on availability:
/// - Claude models default to Ohio (us-east5) /// - Claude models default to Ohio (us-east5)
/// - Gemini models default to Iowa (us-central1) /// - Gemini models default to Iowa (us-central1)
/// - MaaS models default to Iowa (us-central1)
pub fn known_location(&self) -> GcpLocation { pub fn known_location(&self) -> GcpLocation {
match self { match self {
Self::Claude(_) => GcpLocation::Ohio, Self::Claude(_) => GcpLocation::Ohio,
Self::Gemini(_) => GcpLocation::Iowa, Self::Gemini(_) => GcpLocation::Iowa,
Self::MaaS(_, _) => GcpLocation::Iowa,
} }
} }
} }
@@ -162,6 +168,15 @@ impl TryFrom<&str> for GcpVertexAIModel {
"gemini-2.5-pro-preview-05-06" => Ok(Self::Gemini(GeminiVersion::Pro25Preview)), "gemini-2.5-pro-preview-05-06" => Ok(Self::Gemini(GeminiVersion::Pro25Preview)),
"gemini-2.5-flash" => Ok(Self::Gemini(GeminiVersion::Flash25)), "gemini-2.5-flash" => Ok(Self::Gemini(GeminiVersion::Flash25)),
"gemini-2.5-pro" => Ok(Self::Gemini(GeminiVersion::Pro25)), "gemini-2.5-pro" => Ok(Self::Gemini(GeminiVersion::Pro25)),
// MaaS models (Model as a Service from Model Garden)
_ if s.ends_with("-maas") => {
let publisher = s
.split('-')
.next()
.ok_or_else(|| ModelError::UnsupportedModel(s.to_string()))?
.to_string();
Ok(Self::MaaS(publisher, s.to_string()))
}
// Generic models based on prefix matching // Generic models based on prefix matching
_ if s.starts_with("claude-") => { _ if s.starts_with("claude-") => {
Ok(Self::Claude(ClaudeVersion::Generic(s.to_string()))) Ok(Self::Claude(ClaudeVersion::Generic(s.to_string())))
@@ -202,28 +217,32 @@ impl RequestContext {
/// Returns the provider associated with the model. /// Returns the provider associated with the model.
pub fn provider(&self) -> ModelProvider { pub fn provider(&self) -> ModelProvider {
match self.model { match &self.model {
GcpVertexAIModel::Claude(_) => ModelProvider::Anthropic, GcpVertexAIModel::Claude(_) => ModelProvider::Anthropic,
GcpVertexAIModel::Gemini(_) => ModelProvider::Google, GcpVertexAIModel::Gemini(_) => ModelProvider::Google,
GcpVertexAIModel::MaaS(publisher, _) => ModelProvider::MaaS(publisher.clone()),
} }
} }
} }
/// Represents available model providers. /// Represents available model providers.
#[derive(Debug, Clone, PartialEq, Eq, Copy)] #[derive(Debug, Clone, PartialEq, Eq)]
pub enum ModelProvider { pub enum ModelProvider {
/// Anthropic provider (Claude models) /// Anthropic provider (Claude models)
Anthropic, Anthropic,
/// Google provider (Gemini models) /// Google provider (Gemini models)
Google, Google,
/// MaaS provider (Model as a Service from Model Garden)
MaaS(String),
} }
impl ModelProvider { impl ModelProvider {
/// Returns the string representation of the provider. /// Returns the string representation of the provider.
pub fn as_str(&self) -> &'static str { pub fn as_str(&self) -> String {
match self { match self {
Self::Anthropic => "anthropic", Self::Anthropic => "anthropic".to_string(),
Self::Google => "google", Self::Google => "google".to_string(),
Self::MaaS(publisher) => publisher.clone(),
} }
} }
} }
@@ -305,6 +324,12 @@ pub fn create_request(
GcpVertexAIModel::Gemini(_) => { GcpVertexAIModel::Gemini(_) => {
create_google_request(model_config, system, messages, tools)? create_google_request(model_config, system, messages, tools)?
} }
GcpVertexAIModel::MaaS(_, _) => {
// TODO: Branch on publisher for format selection once we know which
// MaaS providers use which formats (e.g., OpenAI vs Google format)
// For now, default to Google format since most use generateContent endpoint
create_google_request(model_config, system, messages, tools)?
}
}; };
Ok((request, context)) Ok((request, context))
@@ -322,6 +347,7 @@ pub fn response_to_message(response: Value, request_context: RequestContext) ->
match request_context.provider() { match request_context.provider() {
ModelProvider::Anthropic => anthropic::response_to_message(&response), ModelProvider::Anthropic => anthropic::response_to_message(&response),
ModelProvider::Google => google::response_to_message(response), ModelProvider::Google => google::response_to_message(response),
ModelProvider::MaaS(_) => google::response_to_message(response),
} }
} }
@@ -337,6 +363,7 @@ pub fn get_usage(data: &Value, request_context: &RequestContext) -> Result<Usage
match request_context.provider() { match request_context.provider() {
ModelProvider::Anthropic => anthropic::get_usage(data), ModelProvider::Anthropic => anthropic::get_usage(data),
ModelProvider::Google => google::get_usage(data), ModelProvider::Google => google::get_usage(data),
ModelProvider::MaaS(_) => google::get_usage(data),
} }
} }
+7 -2
View File
@@ -197,6 +197,7 @@ impl GcpVertexAIProvider {
let endpoint = match provider { let endpoint = match provider {
ModelProvider::Anthropic => "streamRawPredict", ModelProvider::Anthropic => "streamRawPredict",
ModelProvider::Google => "generateContent", ModelProvider::Google => "generateContent",
ModelProvider::MaaS(_) => "generateContent",
}; };
// Construct path for URL // Construct path for URL
@@ -586,8 +587,12 @@ mod tests {
#[test] #[test]
fn test_model_provider_conversion() { fn test_model_provider_conversion() {
assert_eq!(ModelProvider::Anthropic.as_str(), "anthropic"); assert_eq!(ModelProvider::Anthropic.as_str(), "anthropic".to_string());
assert_eq!(ModelProvider::Google.as_str(), "google"); assert_eq!(ModelProvider::Google.as_str(), "google".to_string());
assert_eq!(
ModelProvider::MaaS("qwen".to_string()).as_str(),
"qwen".to_string()
);
} }
#[test] #[test]