diff --git a/crates/goose-provider-types/src/canonical.rs b/crates/goose-provider-types/src/canonical.rs index 5ac718dae..1b05b09fe 100644 --- a/crates/goose-provider-types/src/canonical.rs +++ b/crates/goose-provider-types/src/canonical.rs @@ -65,10 +65,10 @@ pub fn recommended_models_from_registry(provider: &str) -> Vec { .collect() } -/// Providers that run models locally — their cost is always zero regardless -/// of what the canonical registry says for the underlying model architecture. -fn is_local_provider(provider: &str) -> bool { - matches!(provider, "ollama" | "local") +/// Catalog pricing is not valid for local inference or Azure Foundry deployments. +/// Azure billing depends on deployment region, SKU, offer, and contract. +fn should_clear_catalog_pricing(provider: &str) -> bool { + matches!(provider, "ollama" | "local" | "azure_foundry") } pub fn maybe_get_canonical_model(provider: &str, model: &str) -> Option { @@ -81,9 +81,7 @@ pub fn maybe_get_canonical_model(provider: &str, model: &str) -> Option String { fn is_meta_provider(provider: &str) -> bool { matches!( provider, - "databricks" | "databricks_v2" | "tetrate" | "bedrock" | "azure" + "databricks" | "databricks_v2" | "tetrate" | "bedrock" | "azure" | "azure_foundry" ) } @@ -45,7 +45,7 @@ pub fn map_provider_name(provider: &str) -> &str { match provider { // Goose provider names that differ from models.dev names "xai" => "x-ai", - "azure_openai" => "azure", + "azure_openai" | "azure_foundry" => "azure", "aws_bedrock" => "amazon-bedrock", "gcp_vertex_ai" => "google-vertex", "gemini_oauth" => "google", @@ -420,6 +420,10 @@ mod tests { map_to_canonical_model("azure", "gpt-4o", r), Some("openai/gpt-4o".to_string()) ); + assert_eq!( + map_to_canonical_model("azure_foundry", "gpt-4o", r), + Some("openai/gpt-4o".to_string()) + ); // === OpenAI O-series === assert_eq!( diff --git a/crates/goose-provider-types/src/formats/anthropic.rs b/crates/goose-provider-types/src/formats/anthropic.rs index 49484dea7..e114319ed 100644 --- a/crates/goose-provider-types/src/formats/anthropic.rs +++ b/crates/goose-provider-types/src/formats/anthropic.rs @@ -776,6 +776,26 @@ pub fn create_request( messages: &[Message], tools: &[Tool], options: AnthropicFormatOptions, +) -> Result { + create_request_for_model( + provider_name, + model_config, + &model_config.model_name, + system, + messages, + tools, + options, + ) +} + +pub fn create_request_for_model( + provider_name: &str, + model_config: &ModelConfig, + wire_model_name: &str, + system: &str, + messages: &[Message], + tools: &[Tool], + options: AnthropicFormatOptions, ) -> Result { let options = options.for_model(model_config); let anthropic_messages = format_messages_with_options(messages, options); @@ -788,7 +808,7 @@ pub fn create_request( let max_tokens = model_config.max_output_tokens(); let mut payload = json!({ - "model": model_config.model_name, + "model": wire_model_name, "messages": anthropic_messages, "max_tokens": max_tokens, }); diff --git a/crates/goose-provider-types/src/formats/openai.rs b/crates/goose-provider-types/src/formats/openai.rs index b3acacfa1..fbe00f039 100644 --- a/crates/goose-provider-types/src/formats/openai.rs +++ b/crates/goose-provider-types/src/formats/openai.rs @@ -1401,6 +1401,32 @@ pub fn create_request_with_options( image_format: &ImageFormat, for_streaming: bool, format_options: OpenAiFormatOptions, +) -> anyhow::Result { + let (wire_model_name, _) = extract_reasoning_effort(&model_config.model_name); + create_request_for_model_with_options( + model_config, + &wire_model_name, + &model_config.model_name, + system, + messages, + tools, + image_format, + for_streaming, + format_options, + ) +} + +#[allow(clippy::too_many_arguments)] +pub fn create_request_for_model_with_options( + model_config: &ModelConfig, + wire_model_name: &str, + capability_model_name: &str, + system: &str, + messages: &[Message], + tools: &[Tool], + image_format: &ImageFormat, + for_streaming: bool, + format_options: OpenAiFormatOptions, ) -> anyhow::Result { if model_config.model_name.starts_with("o1-mini") { return Err(anyhow!( @@ -1408,7 +1434,7 @@ pub fn create_request_with_options( )); } - let (model_name, legacy_reasoning_effort) = extract_reasoning_effort(&model_config.model_name); + let (model_name, legacy_reasoning_effort) = extract_reasoning_effort(capability_model_name); let is_reasoning_model = is_openai_responses_model(&model_name); let supports_xai_effort = supports_xai_reasoning_effort(&model_name); let reasoning_effort = if is_reasoning_model { @@ -1439,7 +1465,7 @@ pub fn create_request_with_options( messages_array.extend(messages_spec); let mut payload = json!({ - "model": model_name, + "model": wire_model_name, "messages": messages_array }); diff --git a/crates/goose-provider-types/src/formats/openai_responses.rs b/crates/goose-provider-types/src/formats/openai_responses.rs index 32d655657..a4facaf96 100644 --- a/crates/goose-provider-types/src/formats/openai_responses.rs +++ b/crates/goose-provider-types/src/formats/openai_responses.rs @@ -388,6 +388,7 @@ fn add_message_items(input_items: &mut Vec, messages: &[Message]) { MessageContentBlock::ToolRequest(request) if message.role == Role::Assistant => { if !text_items.is_empty() { input_items.push(json!({ + "type": "message", "role": role, "content": text_items })); @@ -435,6 +436,7 @@ fn add_message_items(input_items: &mut Vec, messages: &[Message]) { MessageContentBlock::ToolResponse(response) => { if !text_items.is_empty() { input_items.push(json!({ + "type": "message", "role": role, "content": text_items })); @@ -520,6 +522,7 @@ fn add_message_items(input_items: &mut Vec, messages: &[Message]) { MessageContentBlock::FrontendToolRequest(request) => { if !text_items.is_empty() { input_items.push(json!({ + "type": "message", "role": role, "content": text_items })); @@ -559,6 +562,7 @@ fn add_message_items(input_items: &mut Vec, messages: &[Message]) { if !text_items.is_empty() { input_items.push(json!({ + "type": "message", "role": role, "content": text_items })); @@ -579,11 +583,31 @@ pub fn create_responses_request( system: &str, messages: &[Message], tools: &[Tool], +) -> anyhow::Result { + let (wire_model_name, _) = extract_reasoning_effort(&model_config.model_name); + create_responses_request_for_model( + model_config, + &wire_model_name, + &model_config.model_name, + system, + messages, + tools, + ) +} + +pub fn create_responses_request_for_model( + model_config: &ModelConfig, + wire_model_name: &str, + capability_model_name: &str, + system: &str, + messages: &[Message], + tools: &[Tool], ) -> anyhow::Result { let mut input_items = Vec::new(); if !system.is_empty() { input_items.push(json!({ + "type": "message", "role": "system", "content": [{ "type": "input_text", @@ -594,7 +618,7 @@ pub fn create_responses_request( add_message_items(&mut input_items, messages); - let (model_name, legacy_reasoning_effort) = extract_reasoning_effort(&model_config.model_name); + let (model_name, legacy_reasoning_effort) = extract_reasoning_effort(capability_model_name); // All models routed here are responses-capable; temperature is rejected // by the API for reasoning models regardless of whether an explicit // effort suffix was provided. @@ -639,7 +663,7 @@ pub fn create_responses_request( )); } let mut payload = json!({ - "model": model_name, + "model": wire_model_name, "input": input_items, "store": store, }); @@ -1358,16 +1382,12 @@ mod tests { let types: Vec<&str> = input .iter() - .map(|item| { - item.get("type") - .and_then(|v| v.as_str()) - .unwrap_or_else(|| item["role"].as_str().unwrap()) - }) + .map(|item| item["type"].as_str().unwrap()) .collect(); assert_eq!( types, - vec!["assistant", "function_call", "assistant", "function_call"] + vec!["message", "function_call", "message", "function_call"] ); } diff --git a/crates/goose-providers/src/anthropic.rs b/crates/goose-providers/src/anthropic.rs index 5c8b0e50b..f5248fa79 100644 --- a/crates/goose-providers/src/anthropic.rs +++ b/crates/goose-providers/src/anthropic.rs @@ -16,7 +16,8 @@ use tokio_util::io::StreamReader; use super::api_client::ApiClient; use super::base::{ConfigKey, MessageStream, ModelInfo, Provider, ProviderMetadata}; use super::formats::anthropic::{ - create_request, response_to_streaming_message, AnthropicFormatOptions, ANTHROPIC_PROVIDER_NAME, + create_request_for_model, response_to_streaming_message, AnthropicFormatOptions, + ANTHROPIC_PROVIDER_NAME, }; use super::openai_compatible::handle_status; use super::openai_compatible::map_http_error_to_provider_error; @@ -155,6 +156,54 @@ impl AnthropicProviderBuilder { } impl AnthropicProvider { + pub async fn stream_for_model( + &self, + model_config: &ModelConfig, + wire_model: &str, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result { + let mut payload = create_request_for_model( + ANTHROPIC_PROVIDER_NAME, + model_config, + wire_model, + system, + messages, + tools, + self.format_options, + )?; + payload["stream"] = Value::Bool(true); + let mut log = start_log(model_config, &payload)?; + let response = self + .with_retry(|| async { + handle_status( + self.api_client + .request("v1/messages") + .model_headers(model_config)? + .response_post(&payload) + .await?, + ) + .await + }) + .await + .inspect_err(|e| { + let _ = log.error(e); + })?; + let stream = response.bytes_stream().map_err(io::Error::other); + Ok(Box::pin(try_stream! { + let reader = StreamReader::new(stream); + let framed = tokio_util::codec::FramedRead::new(reader, tokio_util::codec::LinesCodec::new()).map_err(anyhow::Error::from); + let messages = response_to_streaming_message(framed); + pin!(messages); + while let Some(message) = futures::StreamExt::next(&mut messages).await { + let (message, usage) = message.map_err(ProviderError::from_stream_error)?; + log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?; + yield (message, usage); + } + })) + } + async fn fetch_models_from_api(&self) -> Result, ProviderError> { let response = self.api_client.request("v1/models").api_get().await?; @@ -233,6 +282,13 @@ impl Provider for AnthropicProvider { &self.name } + async fn refresh_credentials(&self) -> Result<(), ProviderError> { + self.api_client + .refresh_credentials() + .await + .map_err(|error| ProviderError::Authentication(error.to_string())) + } + fn skip_canonical_filtering(&self) -> bool { self.skip_canonical_filtering } @@ -266,49 +322,14 @@ impl Provider for AnthropicProvider { messages: &[Message], tools: &[Tool], ) -> Result { - let mut payload = create_request( - ANTHROPIC_PROVIDER_NAME, + self.stream_for_model( model_config, + &model_config.model_name, system, messages, tools, - self.format_options, - )?; - payload - .as_object_mut() - .unwrap() - .insert("stream".to_string(), Value::Bool(true)); - - let mut log = start_log(model_config, &payload)?; - - let response = self - .with_retry(|| async { - let request = self - .api_client - .request("v1/messages") - .model_headers(model_config)?; - let resp = request.response_post(&payload).await?; - handle_status(resp).await - }) - .await - .inspect_err(|e| { - let _ = log.error(e); - })?; - - let stream = response.bytes_stream().map_err(io::Error::other); - - Ok(Box::pin(try_stream! { - let stream_reader = StreamReader::new(stream); - let framed = tokio_util::codec::FramedRead::new(stream_reader, tokio_util::codec::LinesCodec::new()).map_err(anyhow::Error::from); - - let message_stream = response_to_streaming_message(framed); - pin!(message_stream); - while let Some(message) = futures::StreamExt::next(&mut message_stream).await { - let (message, usage) = message.map_err(ProviderError::from_stream_error)?; - log.write(&message, usage.as_ref().map(|f| f.usage).as_ref())?; - yield (message, usage); - } - })) + ) + .await } } diff --git a/crates/goose-providers/src/api_client.rs b/crates/goose-providers/src/api_client.rs index 57a44fc8c..5ba3b2a30 100644 --- a/crates/goose-providers/src/api_client.rs +++ b/crates/goose-providers/src/api_client.rs @@ -195,6 +195,10 @@ fn convert_key_to_pkcs8_pem(key_pem_str: &str) -> Result { #[async_trait] pub trait AuthProvider: Send + Sync { async fn get_auth_header(&self) -> Result<(String, String)>; + + async fn refresh_credentials(&self) -> Result<()> { + anyhow::bail!("credential refresh not supported") + } } pub struct ApiResponse { @@ -356,6 +360,13 @@ impl ApiClient { } } + pub async fn refresh_credentials(&self) -> Result<()> { + match &self.auth { + AuthMethod::Custom(provider) => provider.refresh_credentials().await, + _ => anyhow::bail!("credential refresh not supported"), + } + } + pub async fn api_post(&self, path: &str, payload: &Value) -> Result { self.request(path).api_post(payload).await } diff --git a/crates/goose-providers/src/azure_foundry.rs b/crates/goose-providers/src/azure_foundry.rs new file mode 100644 index 000000000..5eee2808b --- /dev/null +++ b/crates/goose-providers/src/azure_foundry.rs @@ -0,0 +1,1175 @@ +use std::collections::HashMap; +use std::sync::Mutex; + +use anyhow::Result; +use async_trait::async_trait; +use rmcp::model::Tool; + +use crate::anthropic::{AnthropicProvider, AnthropicProviderBuilder, ANTHROPIC_API_VERSION}; +use crate::api_client::{ApiClient, AuthMethod, RequestBuilderDecorator, TlsConfig}; +use crate::base::{ + ConfigKey, MessageStream, ModelInfo, Provider, ProviderDescriptor, ProviderMetadata, +}; +use crate::conversation::message::Message; +use crate::errors::ProviderError; +use crate::formats::openai::is_openai_responses_model; +use crate::model::ModelConfig; +use crate::openai::{OpenAiProvider, OpenAiProviderBuilder}; +use crate::openai_compatible::{handle_response_openai_compat, OpenAiCompatibleProvider}; + +pub const AZURE_FOUNDRY_PROVIDER_NAME: &str = "azure_foundry"; +pub const AZURE_FOUNDRY_DEFAULT_MODEL: &str = "Phi-4"; +pub const AZURE_FOUNDRY_DOC_URL: &str = + "https://learn.microsoft.com/azure/ai-foundry/foundry-models/how-to/inference"; + +pub const AZURE_FOUNDRY_KNOWN_MODELS: &[&str] = &[ + "Phi-4", + "Phi-4-mini", + "Meta-Llama-3.3-70B-Instruct", + "Mistral-large-2411", + "Cohere-command-r-plus-08-2024", + "AI21-Jamba-1.5-Large", + "DeepSeek-R1", + "DeepSeek-V3", + "glm-4.7", + "Kimi-K2-Instruct", + "claude-sonnet-4-6", + "claude-opus-4-6", + "gpt-5", + "o3", +]; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum EndpointKind { + Maas, + Resource, + Project, +} + +pub fn endpoint_kind(endpoint: &str) -> EndpointKind { + if endpoint.contains("/api/projects/") { + EndpointKind::Project + } else if endpoint + .split_once("://") + .map(|(_, rest)| rest) + .unwrap_or(endpoint) + .split('/') + .next() + .is_some_and(|host| host.ends_with(".services.ai.azure.com")) + { + EndpointKind::Resource + } else { + EndpointKind::Maas + } +} + +pub fn is_project_endpoint(endpoint: &str) -> bool { + endpoint_kind(endpoint) == EndpointKind::Project +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ModelPublisher { + OpenAi, + Anthropic, + Partner, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct DeploymentMetadata { + publisher: ModelPublisher, + model_name: String, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum InferenceRoute { + MaasChatCompletions, + ProjectChatCompletions, + ProjectResponses, + AnthropicMessages, +} + +fn inference_route( + project_endpoint: bool, + publisher: ModelPublisher, + underlying_model: &str, +) -> InferenceRoute { + if !project_endpoint { + return InferenceRoute::MaasChatCompletions; + } + match publisher { + ModelPublisher::Anthropic => InferenceRoute::AnthropicMessages, + ModelPublisher::OpenAi if is_openai_responses_model(underlying_model) => { + InferenceRoute::ProjectResponses + } + ModelPublisher::OpenAi | ModelPublisher::Partner => InferenceRoute::ProjectChatCompletions, + } +} + +impl ModelPublisher { + fn from_azure(value: &str) -> Self { + match value.to_ascii_lowercase().as_str() { + "openai" => Self::OpenAi, + "anthropic" => Self::Anthropic, + _ => Self::Partner, + } + } + + fn from_model_name(value: &str) -> Self { + let value = value.to_ascii_lowercase(); + if value.starts_with("claude") { + Self::Anthropic + } else if is_openai_responses_model(&value) { + Self::OpenAi + } else { + Self::Partner + } + } +} + +pub struct AzureFoundryProvider { + chat: OpenAiCompatibleProvider, + responses: Option, + anthropic: Option, + deployments_client: ApiClient, + endpoint: String, + api_version: Option, + maas_model: Option, + deployments: Mutex>, +} + +impl ProviderDescriptor for AzureFoundryProvider { + fn metadata() -> ProviderMetadata { + ProviderMetadata::new( + AZURE_FOUNDRY_PROVIDER_NAME, + "Azure AI Foundry", + "OpenAI, Anthropic, and partner models deployed through Azure AI Foundry", + AZURE_FOUNDRY_DEFAULT_MODEL, + AZURE_FOUNDRY_KNOWN_MODELS.to_vec(), + AZURE_FOUNDRY_DOC_URL, + vec![ + ConfigKey::new("AZURE_FOUNDRY_ENDPOINT", true, false, None, true), + ConfigKey::new("AZURE_FOUNDRY_API_KEY", false, true, Some(""), true), + ConfigKey::new("AZURE_FOUNDRY_AD_TOKEN", false, true, Some(""), false), + ConfigKey::new("AZURE_FOUNDRY_MODEL", false, false, None, true), + ConfigKey::new("AZURE_FOUNDRY_API_VERSION", false, false, None, false), + ], + ) + } +} + +impl AzureFoundryProvider { + #[allow(clippy::too_many_arguments)] + pub fn create( + endpoint: String, + api_version: Option, + maas_model: Option, + chat_auth: AuthMethod, + responses_auth: AuthMethod, + anthropic_auth: AuthMethod, + deployments_auth: AuthMethod, + tls_config: Option, + request_builder: Option, + ) -> Result { + let endpoint = endpoint.trim_end_matches('/').to_string(); + let endpoint_kind = endpoint_kind(&endpoint); + let native_inference = endpoint_kind != EndpointKind::Maas; + let maas_model = if native_inference { + None + } else { + Some( + maas_model + .filter(|model| !model.trim().is_empty()) + .ok_or_else(|| { + anyhow::anyhow!("AZURE_FOUNDRY_MODEL is required for MaaS endpoints") + })? + .trim() + .to_string(), + ) + }; + let chat_prefix = if native_inference { + "openai/v1/" + } else { + "v1/" + }; + + let chat_client = configured_client( + endpoint.clone(), + chat_auth, + tls_config.clone(), + request_builder.clone(), + )?; + let chat = OpenAiCompatibleProvider::new( + AZURE_FOUNDRY_PROVIDER_NAME.to_string(), + chat_client, + chat_prefix.to_string(), + ); + + let (responses, anthropic) = if native_inference { + let responses_client = configured_client( + endpoint.clone(), + responses_auth, + tls_config.clone(), + request_builder.clone(), + )?; + let responses = OpenAiProviderBuilder::new(responses_client) + .name(AZURE_FOUNDRY_PROVIDER_NAME) + .base_path("openai/v1/responses") + .skip_canonical_filtering(true) + .build(); + + let hub = endpoint.split("/api/projects/").next().unwrap_or(&endpoint); + let anthropic_client = configured_client( + format!("{hub}/anthropic"), + anthropic_auth, + tls_config.clone(), + request_builder.clone(), + )? + .with_header("anthropic-version", ANTHROPIC_API_VERSION)?; + let anthropic = AnthropicProviderBuilder::new(anthropic_client) + .name(AZURE_FOUNDRY_PROVIDER_NAME) + .skip_canonical_filtering(true) + .build(); + (Some(responses), Some(anthropic)) + } else { + (None, None) + }; + + let deployments_client = configured_client( + endpoint.clone(), + deployments_auth, + tls_config, + request_builder, + )?; + + Ok(Self { + chat, + responses, + anthropic, + deployments_client, + endpoint, + api_version, + maas_model, + deployments: Mutex::new(HashMap::new()), + }) + } + + async fn fetch_deployments( + &self, + ) -> Result<(Vec, HashMap), ProviderError> { + let version = self + .api_version + .as_deref() + .or_else(|| is_project_endpoint(&self.endpoint).then_some("v1")); + let mut next = Some(match version { + Some(version) => format!("deployments?api-version={version}"), + None => "deployments".to_string(), + }); + let mut models = Vec::new(); + let mut deployments = HashMap::new(); + + while let Some(path) = next { + let response = self + .deployments_client + .response_get(&path) + .await + .map_err(|error| ProviderError::NetworkError(error.to_string()))?; + let json = handle_response_openai_compat(response).await?; + if let Some(items) = json.get("value").and_then(|value| value.as_array()) { + for item in items { + let Some(name) = item.get("name").and_then(|value| value.as_str()) else { + continue; + }; + let model_name = item + .get("modelName") + .and_then(|value| value.as_str()) + .unwrap_or(name); + let publisher = item + .get("modelPublisher") + .and_then(|value| value.as_str()) + .map(ModelPublisher::from_azure) + .unwrap_or_else(|| ModelPublisher::from_model_name(model_name)); + models.push(name.to_string()); + deployments.insert( + name.to_string(), + DeploymentMetadata { + publisher, + model_name: model_name.to_string(), + }, + ); + } + } + next = json + .get("nextLink") + .and_then(|value| value.as_str()) + .map(|link| with_api_version(link, version)); + } + + models.sort(); + models.dedup(); + Ok((models, deployments)) + } + + async fn deployment_for(&self, deployment_name: &str) -> Option { + if let Some(deployment) = self + .deployments + .lock() + .expect("Azure Foundry deployment cache poisoned") + .get(deployment_name) + .cloned() + { + return Some(deployment); + } + + let (_, deployments) = self.fetch_deployments().await.ok()?; + let deployment = deployments.get(deployment_name).cloned(); + *self + .deployments + .lock() + .expect("Azure Foundry deployment cache poisoned") = deployments; + deployment + } +} + +fn with_api_version(link: &str, api_version: Option<&str>) -> String { + let Some(api_version) = api_version else { + return link.to_string(); + }; + if link.contains("api-version=") { + return link.to_string(); + } + let separator = if link.contains('?') { '&' } else { '?' }; + format!("{link}{separator}api-version={api_version}") +} + +fn model_info_for_deployment(deployment_name: &str, model_name: &str) -> ModelInfo { + let canonical = crate::canonical::maybe_get_canonical_model("azure_foundry", model_name) + .or_else(|| { + crate::canonical::maybe_get_canonical_model( + "azure_foundry", + &model_name.to_ascii_lowercase(), + ) + }); + ModelInfo { + name: deployment_name.to_string(), + resolved_model: Some(model_name.to_string()), + context_limit: canonical + .as_ref() + .map(|model| model.limit.context) + .unwrap_or_else(|| ModelConfig::new(model_name).context_limit()), + input_token_cost: None, + output_token_cost: None, + currency: None, + supports_cache_control: None, + reasoning: canonical + .and_then(|model| model.reasoning) + .unwrap_or_else(|| ModelConfig::new(model_name).is_reasoning_model()), + } +} + +fn configured_client( + host: String, + auth: AuthMethod, + tls_config: Option, + request_builder: Option, +) -> Result { + let mut client = ApiClient::new_with_tls(host, auth, tls_config)?; + if let Some(request_builder) = request_builder { + client = client.with_request_builder(request_builder); + } + Ok(client) +} + +#[async_trait] +impl Provider for AzureFoundryProvider { + fn get_name(&self) -> &str { + AZURE_FOUNDRY_PROVIDER_NAME + } + + fn skip_canonical_filtering(&self) -> bool { + true + } + + async fn refresh_credentials(&self) -> Result<(), ProviderError> { + self.deployments_client + .refresh_credentials() + .await + .map_err(|error| ProviderError::Authentication(error.to_string())) + } + + async fn fetch_supported_models(&self) -> Result, ProviderError> { + if let Some(model) = &self.maas_model { + return Ok(vec![model.clone()]); + } + if !is_project_endpoint(&self.endpoint) { + return Ok(AZURE_FOUNDRY_KNOWN_MODELS + .iter() + .map(ToString::to_string) + .collect()); + } + let (models, deployments) = self.fetch_deployments().await?; + *self + .deployments + .lock() + .expect("Azure Foundry deployment cache poisoned") = deployments; + Ok(models) + } + + async fn fetch_supported_model_info(&self) -> Result, ProviderError> { + if let Some(model) = &self.maas_model { + return Ok(vec![model_info_for_deployment(model, model)]); + } + if !is_project_endpoint(&self.endpoint) { + return Ok(AZURE_FOUNDRY_KNOWN_MODELS + .iter() + .map(|model| model_info_for_deployment(model, model)) + .collect()); + } + let (models, deployments) = self.fetch_deployments().await?; + let model_info = models + .iter() + .filter_map(|name| { + deployments + .get(name) + .map(|deployment| model_info_for_deployment(name, &deployment.model_name)) + }) + .collect(); + *self + .deployments + .lock() + .expect("Azure Foundry deployment cache poisoned") = deployments; + Ok(model_info) + } + + async fn fetch_model_info(&self, model_name: &str) -> Result { + let resolved_model = if let Some(model) = &self.maas_model { + model.clone() + } else if is_project_endpoint(&self.endpoint) { + self.deployment_for(model_name) + .await + .map(|deployment| deployment.model_name) + .unwrap_or_else(|| model_name.to_string()) + } else { + model_name.to_string() + }; + Ok(model_info_for_deployment(model_name, &resolved_model)) + } + + async fn get_context_limit(&self, model_config: &ModelConfig) -> Result { + if let Some(context_limit) = model_config.context_limit { + return Ok(context_limit); + } + Ok(self + .fetch_model_info(&model_config.model_name) + .await? + .context_limit) + } + + async fn stream( + &self, + model_config: &ModelConfig, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result { + let maas_config = self.maas_model.as_ref().map(|model| { + let mut config = model_config.clone(); + config.model_name = model.clone(); + config + }); + let model_config = maas_config.as_ref().unwrap_or(model_config); + let wire_model = self + .maas_model + .clone() + .unwrap_or_else(|| model_config.model_name.clone()); + let deployment = if is_project_endpoint(&self.endpoint) { + self.deployment_for(&wire_model).await + } else { + None + }; + let publisher = deployment + .as_ref() + .map(|deployment| deployment.publisher) + .unwrap_or_else(|| ModelPublisher::from_model_name(&model_config.model_name)); + let underlying_model = deployment + .as_ref() + .map(|deployment| deployment.model_name.as_str()) + .unwrap_or(&model_config.model_name); + let route = inference_route(self.responses.is_some(), publisher, underlying_model); + let capability_model = deployment + .as_ref() + .map(|deployment| deployment.model_name.as_str()) + .unwrap_or(&model_config.model_name); + let mut capability_config = model_config.clone(); + capability_config.model_name = capability_model.to_string(); + let capability_config = + capability_config.with_canonical_limits(AZURE_FOUNDRY_PROVIDER_NAME); + match route { + InferenceRoute::ProjectResponses => { + self.responses + .as_ref() + .expect("checked above") + .stream_for_model( + &capability_config, + &wire_model, + capability_model, + system, + messages, + tools, + ) + .await + } + InferenceRoute::AnthropicMessages => { + self.anthropic + .as_ref() + .expect("checked above") + .stream_for_model(&capability_config, &wire_model, system, messages, tools) + .await + } + InferenceRoute::MaasChatCompletions | InferenceRoute::ProjectChatCompletions => { + self.chat + .stream_for_model( + &capability_config, + &wire_model, + capability_model, + system, + messages, + tools, + ) + .await + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + use wiremock::matchers::{body_partial_json, method, path, query_param}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + fn project_endpoint(server: &MockServer) -> String { + format!("{}/api/projects/test", server.uri()) + } + + fn raw_model_config(model_name: &str) -> ModelConfig { + let mut config = ModelConfig::new(model_name); + config.model_name = model_name.to_string(); + config + } + + fn project_provider(server: &MockServer) -> AzureFoundryProvider { + AzureFoundryProvider::create( + project_endpoint(server), + None, + None, + AuthMethod::NoAuth, + AuthMethod::NoAuth, + AuthMethod::NoAuth, + AuthMethod::NoAuth, + None, + None, + ) + .unwrap() + } + + fn chat_stream() -> String { + [ + json!({"id":"c1","object":"chat.completion.chunk","model":"test","choices":[{"delta":{"role":"assistant","content":"Hello"},"index":0}]}), + json!({"id":"c1","object":"chat.completion.chunk","choices":[{"delta":{},"finish_reason":"stop","index":0}]}), + json!({"id":"c1","object":"chat.completion.chunk","choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}), + ] + .into_iter() + .map(|chunk| format!("data: {chunk}\n\n")) + .chain(std::iter::once("data: [DONE]\n\n".to_string())) + .collect() + } + + fn responses_stream() -> String { + let created = r#"data: {"type":"response.created","sequence_number":1,"response":{"id":"resp_1","object":"response","created_at":0,"status":"in_progress","model":"gpt-5","output":[]}}"#; + let delta = r#"data: {"type":"response.output_text.delta","sequence_number":2,"item_id":"m1","output_index":0,"content_index":0,"delta":"Hello"}"#; + let completed = r#"data: {"type":"response.completed","sequence_number":3,"response":{"id":"resp_1","object":"response","created_at":0,"status":"completed","model":"gpt-5","output":[],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}}"#; + format!("{created}\n\n{delta}\n\n{completed}\n\ndata: [DONE]\n\n") + } + + fn anthropic_stream() -> String { + let start = r#"data: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-6","stop_reason":null,"usage":{"input_tokens":1,"output_tokens":0}}}"#; + let block_start = r#"data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}"#; + let delta = r#"data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}"#; + let block_stop = r#"data: {"type":"content_block_stop","index":0}"#; + let message_delta = r#"data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":1}}"#; + let stop = r#"data: {"type":"message_stop"}"#; + format!( + "{start}\n\n{block_start}\n\n{delta}\n\n{block_stop}\n\n{message_delta}\n\n{stop}\n\n" + ) + } + + #[test] + fn routing_matrix_uses_endpoint_publisher_and_underlying_model() { + use InferenceRoute::*; + assert_eq!( + inference_route(false, ModelPublisher::OpenAi, "gpt-5"), + MaasChatCompletions + ); + assert_eq!( + inference_route(false, ModelPublisher::Anthropic, "claude-sonnet-4-6"), + MaasChatCompletions + ); + assert_eq!( + inference_route(true, ModelPublisher::OpenAi, "gpt-5"), + ProjectResponses + ); + assert_eq!( + inference_route(true, ModelPublisher::OpenAi, "o3-mini"), + ProjectResponses + ); + assert_eq!( + inference_route(true, ModelPublisher::OpenAi, "gpt-4o"), + ProjectChatCompletions + ); + assert_eq!( + inference_route(true, ModelPublisher::Partner, "gpt-5"), + ProjectChatCompletions + ); + assert_eq!( + inference_route(true, ModelPublisher::Anthropic, "gpt-5"), + AnthropicMessages + ); + } + + #[test] + fn resource_endpoint_routes_declared_models_to_native_surfaces() { + let resource = endpoint_kind("https://hub.services.ai.azure.com"); + let native_inference = resource != EndpointKind::Maas; + + assert_eq!( + inference_route( + native_inference, + ModelPublisher::from_model_name("gpt-5.6-sol"), + "gpt-5.6-sol", + ), + InferenceRoute::ProjectResponses + ); + assert_eq!( + inference_route( + native_inference, + ModelPublisher::from_model_name("claude-sonnet-4-6"), + "claude-sonnet-4-6", + ), + InferenceRoute::AnthropicMessages + ); + assert_eq!( + inference_route( + native_inference, + ModelPublisher::from_model_name("Mistral-large"), + "Mistral-large", + ), + InferenceRoute::ProjectChatCompletions + ); + } + + #[test] + fn model_fallback_only_routes_known_native_families() { + assert_eq!( + ModelPublisher::from_model_name("gpt-5"), + ModelPublisher::OpenAi + ); + assert_eq!( + ModelPublisher::from_model_name("o3-mini"), + ModelPublisher::OpenAi + ); + assert_eq!( + ModelPublisher::from_model_name("claude-sonnet-4-6"), + ModelPublisher::Anthropic + ); + assert_eq!( + ModelPublisher::from_model_name("Mistral-large"), + ModelPublisher::Partner + ); + assert_eq!( + ModelPublisher::from_model_name("glm-4.7"), + ModelPublisher::Partner + ); + assert_eq!( + ModelPublisher::from_model_name("Kimi-K2"), + ModelPublisher::Partner + ); + } + + #[test] + fn endpoint_type_is_detected() { + assert_eq!( + endpoint_kind("https://hub.services.ai.azure.com/api/projects/project"), + EndpointKind::Project + ); + assert_eq!( + endpoint_kind("https://hub.services.ai.azure.com"), + EndpointKind::Resource + ); + assert_eq!( + endpoint_kind("https://deployment.eastus.models.ai.azure.com"), + EndpointKind::Maas + ); + assert!(is_project_endpoint( + "https://hub.services.ai.azure.com/api/projects/project" + )); + assert!(!is_project_endpoint("https://hub.services.ai.azure.com")); + } + + #[test] + fn pagination_link_keeps_api_version() { + assert_eq!( + with_api_version("https://example.test/deployments?page=2", Some("v1")), + "https://example.test/deployments?page=2&api-version=v1" + ); + assert_eq!( + with_api_version( + "https://example.test/deployments?page=2&api-version=2025-05-01", + Some("v1") + ), + "https://example.test/deployments?page=2&api-version=2025-05-01" + ); + } + + #[test] + fn deployment_metadata_enriches_context_without_pricing() { + let info = model_info_for_deployment("production-chat", "gpt-5"); + assert_eq!(info.name, "production-chat"); + assert_eq!(info.resolved_model.as_deref(), Some("gpt-5")); + assert_eq!(info.context_limit, 400_000); + assert_eq!(info.input_token_cost, None); + assert_eq!(info.output_token_cost, None); + } + + #[test] + fn gpt_5_6_sol_uses_its_full_context_window() { + let info = model_info_for_deployment("gpt-5.6-sol", "gpt-5.6-sol"); + + assert_eq!(info.context_limit, 1_050_000); + assert!(info.reasoning); + } + + #[tokio::test] + async fn deployment_discovery_is_paginated_and_preserves_underlying_model() { + let server = MockServer::start().await; + let page_two = format!("{}/api/projects/test/deployments?page=2", server.uri()); + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .and(query_param("api-version", "v1")) + .and(query_param("page", "2")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "value": [{ + "type": "ModelDeployment", + "name": "partner-prod", + "modelName": "Mistral-large", + "modelVersion": "1", + "modelPublisher": "MistralAI" + }] + }))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .and(query_param("api-version", "v1")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "value": [{ + "type": "ModelDeployment", + "name": "openai-prod", + "modelName": "gpt-5", + "modelVersion": "1", + "modelPublisher": "OpenAI" + }], + "nextLink": page_two + }))) + .expect(1) + .mount(&server) + .await; + + let provider = project_provider(&server); + let (names, deployments) = provider.fetch_deployments().await.unwrap(); + assert_eq!(names, vec!["openai-prod", "partner-prod"]); + assert_eq!(deployments["openai-prod"].model_name, "gpt-5"); + assert_eq!(deployments["openai-prod"].publisher, ModelPublisher::OpenAi); + } + + #[tokio::test] + async fn custom_deployment_context_uses_underlying_model() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "value": [{ + "type": "ModelDeployment", + "name": "production-chat", + "modelName": "gpt-5", + "modelVersion": "1", + "modelPublisher": "OpenAI" + }] + }))) + .expect(1) + .mount(&server) + .await; + + let provider = project_provider(&server); + assert_eq!( + provider + .get_context_limit(&ModelConfig::new("production-chat")) + .await + .unwrap(), + 400_000 + ); + } + + #[tokio::test] + async fn canonical_like_alias_context_uses_underlying_model() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "value": [{ + "type": "ModelDeployment", + "name": "gpt-5-high", + "modelName": "custom-128k-model", + "modelPublisher": "OpenAI" + }] + }))) + .expect(1) + .mount(&server) + .await; + + let provider = project_provider(&server); + let config = raw_model_config("gpt-5-high"); + assert_eq!(provider.get_context_limit(&config).await.unwrap(), 128_000); + } + + #[tokio::test] + async fn explicit_context_limit_overrides_deployment_metadata() { + let server = MockServer::start().await; + let provider = project_provider(&server); + let config = raw_model_config("gpt-5-high").with_context_limit(Some(64_000)); + + assert_eq!(provider.get_context_limit(&config).await.unwrap(), 64_000); + assert!(server.received_requests().await.unwrap().is_empty()); + } + + #[tokio::test] + async fn project_inventory_failure_is_propagated() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .respond_with(ResponseTemplate::new(401).set_body_json(json!({ + "error": {"message": "expired"} + }))) + .mount(&server) + .await; + + let provider = project_provider(&server); + assert!(provider.fetch_supported_models().await.is_err()); + assert!(provider.fetch_supported_model_info().await.is_err()); + } + + #[tokio::test] + async fn custom_openai_deployment_uses_alias_and_underlying_capabilities() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "value": [{ + "type": "ModelDeployment", + "name": "production-chat", + "modelName": "gpt-5", + "modelPublisher": "OpenAI" + }] + }))) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/api/projects/test/openai/v1/responses")) + .and(body_partial_json(json!({"model": "production-chat"}))) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(responses_stream()) + .append_header("content-type", "text/event-stream"), + ) + .expect(1) + .mount(&server) + .await; + + let provider = project_provider(&server); + let config = ModelConfig::new("production-chat").with_temperature(Some(0.7)); + provider + .complete(&config, "system", &[], &[]) + .await + .unwrap(); + let requests = server.received_requests().await.unwrap(); + let payload: serde_json::Value = requests + .iter() + .find(|request| request.url.path().ends_with("/responses")) + .unwrap() + .body_json() + .unwrap(); + assert_eq!(payload["model"], "production-chat"); + assert!(payload.get("temperature").is_none()); + assert!(payload.get("capability_model").is_none()); + assert!(config.request_params.is_none()); + assert!(serde_json::to_value(&config) + .unwrap() + .get("capability_model") + .is_none()); + } + + #[tokio::test] + async fn suffixed_deployment_without_metadata_is_preserved_on_the_wire() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"value": []}))) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/api/projects/test/openai/v1/responses")) + .and(body_partial_json(json!({"model": "gpt-5-high"}))) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(responses_stream()) + .append_header("content-type", "text/event-stream"), + ) + .expect(1) + .mount(&server) + .await; + + project_provider(&server) + .complete(&raw_model_config("gpt-5-high"), "system", &[], &[]) + .await + .unwrap(); + } + + #[tokio::test] + async fn suffixed_deployment_alias_is_preserved_on_the_wire() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "value": [{ + "type": "ModelDeployment", + "name": "gpt-5-high", + "modelName": "gpt-5", + "modelPublisher": "OpenAI" + }] + }))) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/api/projects/test/openai/v1/responses")) + .and(body_partial_json(json!({"model": "gpt-5-high"}))) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(responses_stream()) + .append_header("content-type", "text/event-stream"), + ) + .expect(1) + .mount(&server) + .await; + + let provider = project_provider(&server); + let config = raw_model_config("gpt-5-high") + .with_thinking_effort(crate::thinking::ThinkingEffort::Off); + assert_eq!(config.model_name, "gpt-5-high"); + provider + .complete(&config, "system", &[], &[]) + .await + .unwrap(); + } + + #[tokio::test] + async fn anthropic_alias_uses_underlying_output_limit_and_preserves_override() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "value": [{ + "type": "ModelDeployment", + "name": "claude-prod", + "modelName": "claude-sonnet-4-6", + "modelPublisher": "Anthropic" + }] + }))) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/anthropic/v1/messages")) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(anthropic_stream()) + .append_header("content-type", "text/event-stream"), + ) + .expect(2) + .mount(&server) + .await; + + let provider = project_provider(&server); + provider + .complete(&ModelConfig::new("claude-prod"), "system", &[], &[]) + .await + .unwrap(); + provider + .complete( + &ModelConfig::new("claude-prod").with_max_tokens(Some(12_345)), + "system", + &[], + &[], + ) + .await + .unwrap(); + + let requests = server.received_requests().await.unwrap(); + let max_tokens: Vec = requests + .iter() + .filter(|request| request.url.path().ends_with("/messages")) + .map(|request| { + request.body_json::().unwrap()["max_tokens"] + .as_i64() + .unwrap() + }) + .collect(); + assert_eq!(max_tokens, vec![128_000, 12_345]); + } + + #[tokio::test] + async fn canonical_claude_deployment_uses_canonical_output_limit() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "value": [{ + "type": "ModelDeployment", + "name": "claude-sonnet-4-6", + "modelName": "claude-sonnet-4-6", + "modelPublisher": "Anthropic" + }] + }))) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/anthropic/v1/messages")) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(anthropic_stream()) + .append_header("content-type", "text/event-stream"), + ) + .expect(1) + .mount(&server) + .await; + + let provider = project_provider(&server); + provider + .complete(&raw_model_config("claude-sonnet-4-6"), "system", &[], &[]) + .await + .unwrap(); + + let request = server + .received_requests() + .await + .unwrap() + .into_iter() + .find(|request| request.url.path().ends_with("/messages")) + .unwrap(); + let payload = request.body_json::().unwrap(); + assert_eq!(payload["max_tokens"], 128_000); + } + + #[tokio::test] + async fn maas_uses_v1_chat_completions_path_and_bound_model() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/chat/completions")) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(chat_stream()) + .append_header("content-type", "text/event-stream"), + ) + .expect(1) + .mount(&server) + .await; + let provider = AzureFoundryProvider::create( + server.uri(), + None, + Some("bound-model".to_string()), + AuthMethod::NoAuth, + AuthMethod::NoAuth, + AuthMethod::NoAuth, + AuthMethod::NoAuth, + None, + None, + ) + .unwrap(); + let config = raw_model_config("wrong-model"); + provider + .complete(&config, "system", &[], &[]) + .await + .unwrap(); + let request = server.received_requests().await.unwrap().pop().unwrap(); + let payload = request.body_json::().unwrap(); + assert_eq!(payload["model"], "bound-model"); + assert!(payload.get("azure_foundry_deployment").is_none()); + assert_eq!( + provider.fetch_supported_models().await.unwrap(), + vec!["bound-model"] + ); + } + + #[tokio::test] + async fn custom_deployment_names_route_to_all_three_surfaces() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/projects/test/deployments")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "value": [ + {"type":"ModelDeployment","name":"openai-prod","modelName":"gpt-5","modelVersion":"1","modelPublisher":"OpenAI"}, + {"type":"ModelDeployment","name":"claude-prod","modelName":"claude-sonnet-4-6","modelVersion":"1","modelPublisher":"Anthropic"}, + {"type":"ModelDeployment","name":"partner-prod","modelName":"Mistral-large","modelVersion":"1","modelPublisher":"MistralAI"} + ] + }))) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/api/projects/test/openai/v1/responses")) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(responses_stream()) + .append_header("content-type", "text/event-stream"), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/anthropic/v1/messages")) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(anthropic_stream()) + .append_header("content-type", "text/event-stream"), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/api/projects/test/openai/v1/chat/completions")) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(chat_stream()) + .append_header("content-type", "text/event-stream"), + ) + .expect(1) + .mount(&server) + .await; + + let provider = project_provider(&server); + for deployment in ["openai-prod", "claude-prod", "partner-prod"] { + provider + .complete(&ModelConfig::new(deployment), "system", &[], &[]) + .await + .unwrap(); + } + } +} diff --git a/crates/goose-providers/src/lib.rs b/crates/goose-providers/src/lib.rs index 77e011b17..745426458 100644 --- a/crates/goose-providers/src/lib.rs +++ b/crates/goose-providers/src/lib.rs @@ -1,5 +1,6 @@ pub mod anthropic; pub mod api_client; +pub mod azure_foundry; pub mod databricks; pub mod databricks_auth; pub mod databricks_v2; diff --git a/crates/goose-providers/src/openai.rs b/crates/goose-providers/src/openai.rs index e033edaef..d97c30aee 100644 --- a/crates/goose-providers/src/openai.rs +++ b/crates/goose-providers/src/openai.rs @@ -11,7 +11,8 @@ use crate::formats::openai::{ create_request_with_options, get_cost, get_usage, response_to_message, OpenAiFormatOptions, }; use crate::formats::openai_responses::{ - create_responses_request, get_responses_usage, responses_api_to_message, ResponsesApiResponse, + create_responses_request_for_model, get_responses_usage, responses_api_to_message, + ResponsesApiResponse, }; use crate::images::ImageFormat; use crate::openai_compatible::{ @@ -272,6 +273,80 @@ impl OpenAiProviderBuilder { } impl OpenAiProvider { + pub async fn stream_for_model( + &self, + model_config: &ModelConfig, + wire_model: &str, + capability_model: &str, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result { + let mut payload = create_responses_request_for_model( + model_config, + wire_model, + capability_model, + system, + messages, + tools, + )?; + payload["stream"] = serde_json::Value::Bool(self.supports_streaming); + self.stream_responses_payload(model_config, payload).await + } + + async fn stream_responses_payload( + &self, + model_config: &ModelConfig, + payload: serde_json::Value, + ) -> Result { + let mut log = start_log(model_config, &payload)?; + let response = self + .with_retry(|| async { + handle_status( + self.api_client + .request(&Self::map_base_path( + &self.base_path, + "responses", + OPEN_AI_DEFAULT_RESPONSES_PATH, + )) + .model_headers(model_config)? + .response_post(&payload) + .await?, + ) + .await + }) + .await + .inspect_err(|e| { + let _ = log.error(e); + })?; + if self.supports_streaming { + stream_responses_compat(response, log) + } else { + let json: serde_json::Value = response.json().await.map_err(|e| { + ProviderError::RequestFailed(format!("Failed to parse JSON: {}", e)) + })?; + let parsed: ResponsesApiResponse = + serde_json::from_value(json.clone()).map_err(|e| { + ProviderError::ExecutionError(format!( + "Failed to parse responses API response: {}", + e + )) + })?; + let message = responses_api_to_message(&parsed)?; + let usage_data = get_responses_usage(&parsed); + let usage_json = json.get("usage").unwrap_or(&serde_json::Value::Null); + let mut usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); + if let Some(cost) = get_cost(usage_json) { + usage = usage.with_cost(cost, CostSource::ProviderReported); + } + log.write( + &serde_json::to_value(&message).unwrap_or_default(), + Some(&usage_data), + )?; + Ok(super::base::stream_from_single_message(message, usage)) + } + } + #[doc(hidden)] pub fn new(api_client: ApiClient) -> Self { Self { @@ -584,6 +659,13 @@ impl Provider for OpenAiProvider { &self.name } + async fn refresh_credentials(&self) -> Result<(), ProviderError> { + self.api_client + .refresh_credentials() + .await + .map_err(|error| ProviderError::Authentication(error.to_string())) + } + fn skip_canonical_filtering(&self) -> bool { self.skip_canonical_filtering } @@ -655,61 +737,17 @@ impl Provider for OpenAiProvider { tools: &[Tool], ) -> Result { if self.should_use_responses_api_for_provider(&model_config.model_name) { - let mut payload = create_responses_request(model_config, system, messages, tools)?; - payload["stream"] = serde_json::Value::Bool(self.supports_streaming); - - let mut log = start_log(model_config, &payload)?; - - let response = self - .with_retry(|| async { - let payload_clone = payload.clone(); - let resp = self - .api_client - .request(&Self::map_base_path( - &self.base_path, - "responses", - OPEN_AI_DEFAULT_RESPONSES_PATH, - )) - .model_headers(model_config)? - .response_post(&payload_clone) - .await?; - handle_status(resp).await - }) - .await - .inspect_err(|e| { - let _ = log.error(e); - })?; - - if self.supports_streaming { - stream_responses_compat(response, log) - } else { - let json: serde_json::Value = response.json().await.map_err(|e| { - ProviderError::RequestFailed(format!("Failed to parse JSON: {}", e)) - })?; - - let responses_api_response: ResponsesApiResponse = - serde_json::from_value(json.clone()).map_err(|e| { - ProviderError::ExecutionError(format!( - "Failed to parse responses API response: {}", - e - )) - })?; - - let message = responses_api_to_message(&responses_api_response)?; - let usage_data = get_responses_usage(&responses_api_response); - let usage_json = json.get("usage").unwrap_or(&serde_json::Value::Null); - let mut usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); - if let Some(cost) = get_cost(usage_json) { - usage = usage.with_cost(cost, CostSource::ProviderReported); - } - - log.write( - &serde_json::to_value(&message).unwrap_or_default(), - Some(&usage_data), - )?; - - Ok(super::base::stream_from_single_message(message, usage)) - } + let (wire_model, _) = + crate::formats::openai::extract_reasoning_effort(&model_config.model_name); + self.stream_for_model( + model_config, + &wire_model, + &model_config.model_name, + system, + messages, + tools, + ) + .await } else { let payload = create_request_with_options( model_config, diff --git a/crates/goose-providers/src/openai_compatible.rs b/crates/goose-providers/src/openai_compatible.rs index b9e4648a4..2d0c12772 100644 --- a/crates/goose-providers/src/openai_compatible.rs +++ b/crates/goose-providers/src/openai_compatible.rs @@ -18,7 +18,8 @@ use super::retry::ProviderRetry; use crate::conversation::message::Message; use crate::errors::ProviderError; use crate::formats::openai::{ - create_request, get_cost, get_usage, response_to_message, response_to_streaming_message, + create_request, create_request_for_model_with_options, get_cost, get_usage, + response_to_message, response_to_streaming_message, OpenAiFormatOptions, }; use crate::formats::openai_responses::responses_api_to_streaming_message; use crate::model::ModelConfig; @@ -49,6 +50,99 @@ impl OpenAiCompatibleProvider { self } + #[allow(clippy::too_many_arguments)] + fn build_request_for_model( + &self, + model_config: &ModelConfig, + wire_model: &str, + capability_model: &str, + system: &str, + messages: &[Message], + tools: &[Tool], + for_streaming: bool, + ) -> Result { + create_request_for_model_with_options( + model_config, + wire_model, + capability_model, + system, + messages, + tools, + &ImageFormat::OpenAi, + for_streaming, + OpenAiFormatOptions { + preserve_thinking_context: true, + }, + ) + .map_err(|e| ProviderError::RequestFailed(format!("Failed to create request: {}", e))) + } + + pub async fn stream_for_model( + &self, + model_config: &ModelConfig, + wire_model: &str, + capability_model: &str, + system: &str, + messages: &[Message], + tools: &[Tool], + ) -> Result { + let payload = self.build_request_for_model( + model_config, + wire_model, + capability_model, + system, + messages, + tools, + self.supports_streaming, + )?; + self.stream_payload(model_config, payload).await + } + + async fn stream_payload( + &self, + model_config: &ModelConfig, + payload: Value, + ) -> Result { + let mut log = start_log(model_config, &payload)?; + let path = format!("{}chat/completions", self.completions_prefix); + let response = self + .with_retry(|| async { + handle_status( + self.api_client + .request(&path) + .model_headers(model_config)? + .response_post(&payload) + .await?, + ) + .await + }) + .await + .inspect_err(|e| { + let _ = log.error(e); + })?; + if self.supports_streaming { + stream_openai_compat(response, log) + } else { + let json = response.json().await.map_err(|e| { + ProviderError::RequestFailed(format!("Failed to parse JSON: {}", e)) + })?; + let message = response_to_message(&json).map_err(|e| { + ProviderError::RequestFailed(format!("Failed to parse message: {}", e)) + })?; + let usage_json = json.get("usage").unwrap_or(&Value::Null); + let usage_data = get_usage(usage_json); + let mut usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); + if let Some(cost) = get_cost(usage_json) { + usage = usage.with_cost(cost, CostSource::ProviderReported); + } + log.write( + &serde_json::to_value(&message).unwrap_or_default(), + Some(&usage.usage), + )?; + Ok(stream_from_single_message(message, usage)) + } + } + fn build_request( &self, model_config: &ModelConfig, @@ -75,6 +169,13 @@ impl Provider for OpenAiCompatibleProvider { &self.name } + async fn refresh_credentials(&self) -> Result<(), ProviderError> { + self.api_client + .refresh_credentials() + .await + .map_err(|error| ProviderError::Authentication(error.to_string())) + } + async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client @@ -116,49 +217,7 @@ impl Provider for OpenAiCompatibleProvider { tools, self.supports_streaming, )?; - let mut log = start_log(model_config, &payload)?; - - let completions_path = format!("{}chat/completions", self.completions_prefix); - let response = self - .with_retry(|| async { - let resp = self - .api_client - .request(&completions_path) - .model_headers(model_config)? - .response_post(&payload) - .await?; - handle_status(resp).await - }) - .await - .inspect_err(|e| { - let _ = log.error(e); - })?; - - if self.supports_streaming { - stream_openai_compat(response, log) - } else { - let json: serde_json::Value = response.json().await.map_err(|e| { - ProviderError::RequestFailed(format!("Failed to parse JSON: {}", e)) - })?; - - let message = response_to_message(&json).map_err(|e| { - ProviderError::RequestFailed(format!("Failed to parse message: {}", e)) - })?; - - let usage_json = json.get("usage").unwrap_or(&serde_json::Value::Null); - let usage_data = get_usage(usage_json); - let mut usage = ProviderUsage::new(model_config.model_name.clone(), usage_data); - if let Some(cost) = get_cost(usage_json) { - usage = usage.with_cost(cost, CostSource::ProviderReported); - } - - log.write( - &serde_json::to_value(&message).unwrap_or_default(), - Some(&usage.usage), - )?; - - Ok(stream_from_single_message(message, usage)) - } + self.stream_payload(model_config, payload).await } } diff --git a/crates/goose/src/model_config.rs b/crates/goose/src/model_config.rs index abd1584b9..5fe2d3c7c 100644 --- a/crates/goose/src/model_config.rs +++ b/crates/goose/src/model_config.rs @@ -14,7 +14,7 @@ pub fn model_config_from_user_config( provider_name: &str, model_name: impl AsRef, ) -> Result { - let model = base_model_config_from_user_config(model_name.as_ref())?; + let model = base_model_config_from_user_config(provider_name, model_name.as_ref())?; materialize_model_config(provider_name, model) } @@ -26,18 +26,26 @@ pub fn model_config_from_user_config_with_session_settings( context_limit: Option, ) -> Result { let config = Config::global(); - let model = base_model_config_from_user_config(model_name.as_ref())?; + let model = base_model_config_from_user_config(provider_name, model_name.as_ref())?; let model = materialize_model_config_inner(model, provider_name, false)? .with_context_limit(context_limit) .with_inherited_session_settings_from(previous, request_params) .with_default_thinking_effort(config.get_goose_thinking_effort()); - Ok(model.with_canonical_limits(provider_name)) + Ok(apply_canonical_limits(provider_name, model)) } pub fn materialize_model_config(provider_name: &str, model: ModelConfig) -> Result { let model = materialize_model_config_inner(model, provider_name, true)?; - Ok(model.with_canonical_limits(provider_name)) + Ok(apply_canonical_limits(provider_name, model)) +} + +fn apply_canonical_limits(provider_name: &str, model: ModelConfig) -> ModelConfig { + if provider_name == goose_providers::azure_foundry::AZURE_FOUNDRY_PROVIDER_NAME { + model + } else { + model.with_canonical_limits(provider_name) + } } fn materialize_model_config_inner( @@ -169,7 +177,10 @@ fn apply_openai_request_params(mut model: ModelConfig) -> ModelConfig { model } -fn base_model_config_from_user_config(model_name: &str) -> Result { +fn base_model_config_from_user_config( + provider_name: &str, + model_name: &str, +) -> Result { let config = Config::global(); let mut model = ModelConfig { model_name: model_name.to_string(), @@ -182,7 +193,9 @@ fn base_model_config_from_user_config(model_name: &str) -> Result { reasoning: None, request_headers: None, }; - model.normalize_effort_suffix(); + if provider_name != goose_providers::azure_foundry::AZURE_FOUNDRY_PROVIDER_NAME { + model.normalize_effort_suffix(); + } Ok(model) } @@ -247,3 +260,27 @@ fn parse_yaml_bool_config(key: &str, value: serde_yaml::Value) -> Result { } } } + +#[cfg(test)] +mod azure_foundry_tests { + use super::*; + + #[test] + fn deployment_name_survives_thinking_effort_changes() { + let config = base_model_config_from_user_config("azure_foundry", "gpt-5-high") + .unwrap() + .with_thinking_effort(ThinkingEffort::Off); + + assert_eq!(config.model_name, "gpt-5-high"); + assert_eq!(config.context_limit, None); + assert_eq!(config.thinking_effort(), Some(ThinkingEffort::Off)); + } + + #[test] + fn none_suffixed_deployment_name_is_preserved() { + let config = base_model_config_from_user_config("azure_foundry", "gpt-5-none").unwrap(); + + assert_eq!(config.model_name, "gpt-5-none"); + assert_eq!(config.thinking_effort(), None); + } +} diff --git a/crates/goose/src/providers/azure_foundry_def.rs b/crates/goose/src/providers/azure_foundry_def.rs new file mode 100644 index 000000000..ad961059f --- /dev/null +++ b/crates/goose/src/providers/azure_foundry_def.rs @@ -0,0 +1,173 @@ +use std::sync::Arc; + +use anyhow::Result; +use async_trait::async_trait; +use futures::future::BoxFuture; +use goose_providers::api_client::{AuthMethod, AuthProvider, TlsConfig}; +use goose_providers::azure_foundry::{endpoint_kind, AzureFoundryProvider, EndpointKind}; +use goose_providers::base::{ProviderDescriptor, ProviderMetadata}; + +use crate::config::{Config, ExtensionConfig}; +use crate::providers::azureauth::{AzureAuth, AzureCredentials}; +use crate::providers::base::ProviderDef; + +const AZURE_PROJECT_ENTRA_RESOURCE: &str = "https://ai.azure.com"; +const AZURE_MAAS_ENTRA_RESOURCE: &str = "https://ml.azure.com"; + +enum AuthHeader { + ApiKey, + Bearer, +} + +struct AzureFoundryAuthProvider { + auth: Arc, + header: AuthHeader, +} + +#[async_trait] +impl AuthProvider for AzureFoundryAuthProvider { + async fn get_auth_header(&self) -> Result<(String, String)> { + let token = self.auth.get_token().await?; + match &self.header { + AuthHeader::ApiKey => Ok(("api-key".to_string(), token.token_value)), + AuthHeader::Bearer => Ok(( + "Authorization".to_string(), + format!("Bearer {}", token.token_value), + )), + } + } + + async fn refresh_credentials(&self) -> Result<()> { + self.auth.invalidate_token().await; + Ok(()) + } +} + +pub struct AzureFoundryProviderDef; + +impl ProviderDescriptor for AzureFoundryProviderDef { + fn metadata() -> ProviderMetadata { + AzureFoundryProvider::metadata() + } +} + +impl ProviderDef for AzureFoundryProviderDef { + type Provider = AzureFoundryProvider; + + fn from_env( + _extensions: Vec, + tls_config: Option, + ) -> BoxFuture<'static, Result> { + Box::pin(from_env(tls_config)) + } +} + +pub async fn from_env(tls_config: Option) -> Result { + let config = Config::global(); + let endpoint: String = config.get_param("AZURE_FOUNDRY_ENDPOINT")?; + let api_version = config.get_param("AZURE_FOUNDRY_API_VERSION").ok(); + let maas_model = config + .get_param::("AZURE_FOUNDRY_MODEL") + .ok() + .filter(|model| !model.trim().is_empty()); + let api_key = config + .get_secret::("AZURE_FOUNDRY_API_KEY") + .ok() + .filter(|key| !key.is_empty()); + let ad_token = config + .get_secret::("AZURE_FOUNDRY_AD_TOKEN") + .ok() + .filter(|token| !token.is_empty()); + let endpoint_kind = endpoint_kind(&endpoint); + let resource = if endpoint_kind == EndpointKind::Maas { + AZURE_MAAS_ENTRA_RESOURCE + } else { + AZURE_PROJECT_ENTRA_RESOURCE + }; + let auth = Arc::new(AzureAuth::new_with_resource( + api_key, + ad_token, + resource.to_string(), + )?); + let auth_method = |header| { + AuthMethod::Custom(Box::new(AzureFoundryAuthProvider { + auth: Arc::clone(&auth), + header, + })) + }; + let anthropic_auth = match auth.credential_type() { + AzureCredentials::ApiKey(key) => AuthMethod::ApiKey { + header_name: "x-api-key".to_string(), + key: key.clone(), + }, + _ => auth_method(AuthHeader::Bearer), + }; + let api_key_auth_header = || match auth.credential_type() { + AzureCredentials::ApiKey(_) => AuthHeader::ApiKey, + _ => AuthHeader::Bearer, + }; + let chat_auth_header = match auth.credential_type() { + AzureCredentials::ApiKey(_) if endpoint_kind == EndpointKind::Maas => AuthHeader::Bearer, + _ => api_key_auth_header(), + }; + + AzureFoundryProvider::create( + endpoint, + api_version, + maas_model, + auth_method(chat_auth_header), + auth_method(api_key_auth_header()), + anthropic_auth, + auth_method(api_key_auth_header()), + tls_config, + Some(crate::session_context::session_id_request_builder()), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + async fn header( + api_key: Option<&str>, + ad_token: Option<&str>, + header: AuthHeader, + ) -> (String, String) { + let auth = Arc::new( + AzureAuth::new_with_resource( + api_key.map(str::to_string), + ad_token.map(str::to_string), + AZURE_PROJECT_ENTRA_RESOURCE.to_string(), + ) + .unwrap(), + ); + AzureFoundryAuthProvider { auth, header } + .get_auth_header() + .await + .unwrap() + } + + #[tokio::test] + async fn project_api_key_uses_api_key_header() { + assert_eq!( + header(Some("key"), None, AuthHeader::ApiKey).await, + ("api-key".to_string(), "key".to_string()) + ); + } + + #[tokio::test] + async fn maas_api_key_uses_bearer_header() { + assert_eq!( + header(Some("key"), None, AuthHeader::Bearer).await, + ("Authorization".to_string(), "Bearer key".to_string()) + ); + } + + #[tokio::test] + async fn entra_token_uses_bearer_header() { + assert_eq!( + header(None, Some("token"), AuthHeader::Bearer).await, + ("Authorization".to_string(), "Bearer token".to_string()) + ); + } +} diff --git a/crates/goose/src/providers/azureauth.rs b/crates/goose/src/providers/azureauth.rs index cce77b585..231fd6563 100644 --- a/crates/goose/src/providers/azureauth.rs +++ b/crates/goose/src/providers/azureauth.rs @@ -59,6 +59,7 @@ struct TokenResponse { #[derive(Debug)] pub struct AzureAuth { credentials: AzureCredentials, + resource: String, cached_token: Arc>>, } @@ -73,6 +74,18 @@ impl AzureAuth { /// # Returns /// * `Result` - A new AzureAuth instance or an error if initialization fails pub fn new(api_key: Option, ad_token: Option) -> Result { + Self::new_with_resource( + api_key, + ad_token, + "https://cognitiveservices.azure.com".to_string(), + ) + } + + pub fn new_with_resource( + api_key: Option, + ad_token: Option, + resource: String, + ) -> Result { let credentials = match (ad_token, api_key) { (Some(token), _) => AzureCredentials::BearerToken(token), (None, Some(key)) => AzureCredentials::ApiKey(key), @@ -81,6 +94,7 @@ impl AzureAuth { Ok(Self { credentials, + resource, cached_token: Arc::new(RwLock::new(None)), }) } @@ -90,6 +104,10 @@ impl AzureAuth { &self.credentials } + pub async fn invalidate_token(&self) { + *self.cached_token.write().await = None; + } + /// Retrieves a valid authentication token. /// /// This method implements an efficient token management strategy: @@ -137,12 +155,7 @@ impl AzureAuth { let az = if cfg!(windows) { "az.cmd" } else { "az" }; let output = tokio::process::Command::new(az) - .args([ - "account", - "get-access-token", - "--resource", - "https://cognitiveservices.azure.com", - ]) + .args(["account", "get-access-token", "--resource", &self.resource]) .set_no_window() .output() .await diff --git a/crates/goose/src/providers/init.rs b/crates/goose/src/providers/init.rs index df88101b8..319681212 100644 --- a/crates/goose/src/providers/init.rs +++ b/crates/goose/src/providers/init.rs @@ -37,6 +37,7 @@ use super::{ }; use crate::config::ExtensionConfig; use crate::providers::anthropic_def::AnthropicProviderDef; +use crate::providers::azure_foundry_def::AzureFoundryProviderDef; use crate::providers::base::ProviderType; use crate::providers::databricks_def::{self, DatabricksProviderDef}; use crate::providers::databricks_v2_def::{self, DatabricksV2ProviderDef}; @@ -69,6 +70,10 @@ async fn init_registry() -> RwLock { ); registry.register::(false); registry.register::(false); + registry.register_with_inventory::( + true, + Some(registrations::azure_foundry_inventory()), + ); #[cfg(feature = "aws-providers")] registry.register::(false); #[cfg(feature = "local-inference")] diff --git a/crates/goose/src/providers/inventory/registrations.rs b/crates/goose/src/providers/inventory/registrations.rs index df7f0ed18..b27679fa6 100644 --- a/crates/goose/src/providers/inventory/registrations.rs +++ b/crates/goose/src/providers/inventory/registrations.rs @@ -20,6 +20,7 @@ use crate::providers::ollama::OLLAMA_PROVIDER_NAME; use crate::providers::openai::{OPEN_AI_DEFAULT_BASE_PATH, OPEN_AI_PROVIDER_NAME}; use crate::providers::pi_acp::{PI_ACP_BINARY, PI_ACP_PROVIDER_NAME}; use crate::providers::xai_oauth::TokenCache as XaiOAuthTokenCache; +use goose_providers::azure_foundry::{endpoint_kind, EndpointKind, AZURE_FOUNDRY_PROVIDER_NAME}; pub fn openai_inventory() -> InventoryRegistration { InventoryRegistration::new(true, || { @@ -67,6 +68,52 @@ pub fn openai_inventory() -> InventoryRegistration { }) } +pub fn azure_foundry_inventory() -> InventoryRegistration { + InventoryRegistration::new(true, || { + let config = Config::global(); + let mut identity = + InventoryIdentityInput::new(AZURE_FOUNDRY_PROVIDER_NAME, AZURE_FOUNDRY_PROVIDER_NAME); + if let Ok(endpoint) = config.get_param::("AZURE_FOUNDRY_ENDPOINT") { + identity = identity.with_public("endpoint", endpoint); + } + if let Ok(api_version) = config.get_param::("AZURE_FOUNDRY_API_VERSION") { + identity = identity.with_public("api_version", api_version); + } + if let Ok(model) = config.get_param::("AZURE_FOUNDRY_MODEL") { + identity = identity.with_public("model", model); + } + if let Some(api_key) = config_secret_value(config, "AZURE_FOUNDRY_API_KEY") { + identity = identity.with_secret("api_key", api_key); + } + if let Some(ad_token) = config_secret_value(config, "AZURE_FOUNDRY_AD_TOKEN") { + identity = identity.with_secret("ad_token", ad_token); + } + Ok(identity) + }) + .with_configured(|| azure_foundry_configured(Config::global())) +} + +fn azure_foundry_configured(config: &Config) -> bool { + azure_foundry_configured_values( + config + .get_param::("AZURE_FOUNDRY_ENDPOINT") + .ok() + .as_deref(), + config + .get_param::("AZURE_FOUNDRY_MODEL") + .ok() + .as_deref(), + ) +} + +fn azure_foundry_configured_values(endpoint: Option<&str>, model: Option<&str>) -> bool { + let Some(endpoint) = endpoint.filter(|endpoint| !endpoint.trim().is_empty()) else { + return false; + }; + endpoint_kind(endpoint) != EndpointKind::Maas + || model.is_some_and(|model| !model.trim().is_empty()) +} + pub fn anthropic_inventory() -> InventoryRegistration { InventoryRegistration::new(true, || { let config = Config::global(); @@ -222,6 +269,30 @@ mod tests { use crate::config::paths::Paths; use chrono::Utc; + #[test] + fn azure_foundry_maas_requires_a_model_to_be_configured() { + assert!(!azure_foundry_configured_values( + Some("https://deployment.models.ai.azure.com"), + None, + )); + assert!(!azure_foundry_configured_values( + Some("https://deployment.models.ai.azure.com"), + Some(" "), + )); + assert!(azure_foundry_configured_values( + Some("https://deployment.models.ai.azure.com"), + Some("Phi-4"), + )); + assert!(azure_foundry_configured_values( + Some("https://hub.services.ai.azure.com/api/projects/project"), + None, + )); + assert!(azure_foundry_configured_values( + Some("https://hub.services.ai.azure.com"), + None, + )); + } + #[test] #[serial_test::serial] fn gemini_oauth_inventory_configured_uses_token_cache() { diff --git a/crates/goose/src/providers/mod.rs b/crates/goose/src/providers/mod.rs index 05fea6beb..11243e5e8 100644 --- a/crates/goose/src/providers/mod.rs +++ b/crates/goose/src/providers/mod.rs @@ -9,6 +9,7 @@ pub mod api_client { } pub mod avian; pub mod azure; +pub mod azure_foundry_def; pub mod azureauth; pub mod base; #[cfg(feature = "aws-providers")] diff --git a/crates/goose/src/session/session_manager.rs b/crates/goose/src/session/session_manager.rs index efbb65a8a..f5160570a 100644 --- a/crates/goose/src/session/session_manager.rs +++ b/crates/goose/src/session/session_manager.rs @@ -715,6 +715,27 @@ impl Session { } } +fn deserialize_session_model_config( + provider_name: Option<&str>, + json: &str, +) -> Option { + let mut model_config: ModelConfig = serde_json::from_str(json).ok()?; + // TODO: Remove this workaround once ModelConfig guarantees deserialize(serialize(config)) == config. + if provider_name == Some(goose_providers::azure_foundry::AZURE_FOUNDRY_PROVIDER_NAME) { + #[derive(Deserialize)] + struct AzurePersistedFields { + model_name: String, + #[serde(default)] + request_params: Option>, + } + + let persisted: AzurePersistedFields = serde_json::from_str(json).ok()?; + model_config.model_name = persisted.model_name; + model_config.request_params = persisted.request_params; + } + Some(model_config) +} + impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session { fn from_row(row: &sqlx::sqlite::SqliteRow) -> Result { use sqlx::Row; @@ -726,8 +747,11 @@ impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session { let user_recipe_values = user_recipe_values_json.and_then(|json| serde_json::from_str(&json).ok()); + let provider_name: Option = row.try_get("provider_name").ok().flatten(); let model_config_json: Option = row.try_get("model_config_json").ok().flatten(); - let model_config = model_config_json.and_then(|json| serde_json::from_str(&json).ok()); + let model_config = model_config_json + .as_deref() + .and_then(|json| deserialize_session_model_config(provider_name.as_deref(), json)); let name: String = { let name_val: String = row.try_get("name").unwrap_or_default(); @@ -788,7 +812,7 @@ impl sqlx::FromRow<'_, sqlx::sqlite::SqliteRow> for Session { conversation: None, message_count: row.try_get("message_count").unwrap_or(0) as usize, last_message_at, - provider_name: row.try_get("provider_name").ok().flatten(), + provider_name, model_config, goose_mode: row .try_get::("goose_mode") @@ -2587,6 +2611,61 @@ mod tests { const NUM_CONCURRENT_SESSIONS: i32 = 10; const GENERATED_SESSION_NAME: &str = "Generated session name"; + #[test] + fn azure_session_model_config_preserves_suffixed_deployment_id() { + let json = serde_json::to_string(&ModelConfig { + model_name: "gpt-5-high".to_string(), + context_limit: None, + temperature: None, + max_tokens: None, + toolshim: false, + toolshim_model: None, + request_params: None, + reasoning: None, + request_headers: None, + }) + .unwrap(); + + let config = deserialize_session_model_config( + Some(goose_providers::azure_foundry::AZURE_FOUNDRY_PROVIDER_NAME), + &json, + ) + .unwrap(); + + assert_eq!(config.model_name, "gpt-5-high"); + assert_eq!(config.thinking_effort(), None); + } + + #[test] + fn azure_session_model_config_preserves_explicit_thinking_effort() { + let config = deserialize_session_model_config( + Some(goose_providers::azure_foundry::AZURE_FOUNDRY_PROVIDER_NAME), + r#"{"model_name":"gpt-5-high","context_limit":null,"temperature":null,"max_tokens":null,"toolshim":false,"toolshim_model":null,"request_params":{"thinking_effort":"low"}}"#, + ) + .unwrap(); + + assert_eq!(config.model_name, "gpt-5-high"); + assert_eq!( + config.thinking_effort(), + Some(goose_providers::thinking::ThinkingEffort::Low) + ); + } + + #[test] + fn non_azure_session_model_config_keeps_suffix_normalization() { + let config = deserialize_session_model_config( + Some(goose_providers::openai::OPEN_AI_PROVIDER_NAME), + r#"{"model_name":"gpt-5-high","context_limit":null,"temperature":null,"max_tokens":null,"toolshim":false,"toolshim_model":null}"#, + ) + .unwrap(); + + assert_eq!(config.model_name, "gpt-5"); + assert_eq!( + config.thinking_effort(), + Some(goose_providers::thinking::ThinkingEffort::High) + ); + } + struct NamingTestProvider; #[async_trait::async_trait] diff --git a/documentation/docs/getting-started/providers.md b/documentation/docs/getting-started/providers.md index ddf72988b..d7200c278 100644 --- a/documentation/docs/getting-started/providers.md +++ b/documentation/docs/getting-started/providers.md @@ -27,6 +27,7 @@ goose is compatible with a wide range of LLM providers, allowing you to choose a | [Anthropic](https://www.anthropic.com/) | Offers Claude, an advanced AI model for natural language tasks. | `ANTHROPIC_API_KEY`, `ANTHROPIC_HOST` (optional) | | [Atomic Chat](https://github.com/AtomicBot-ai/Atomic-Chat) | Run local models with Atomic Chat's OpenAI-compatible server. **Because this provider runs locally, you must first [download a model](#local-llms).** | None required. Connects to local server at `localhost:1337` by default. | | [Avian](https://avian.io/) | Cost-effective inference API with DeepSeek, Kimi, GLM, and MiniMax models. OpenAI-compatible with streaming and function calling support. | `AVIAN_API_KEY`, `AVIAN_HOST` (optional) | +| [Azure AI Foundry](/docs/guides/azure-foundry-provider) | Access OpenAI, Anthropic, Microsoft, Meta, Mistral, DeepSeek, GLM, Kimi, and other models deployed through Azure AI Foundry project or MaaS endpoints. | `AZURE_FOUNDRY_ENDPOINT`, `AZURE_FOUNDRY_API_KEY` (optional), `AZURE_FOUNDRY_AD_TOKEN` (optional), `AZURE_FOUNDRY_API_VERSION` (optional) | | [Azure OpenAI](https://learn.microsoft.com/en-us/azure/ai-services/openai/) | Access Azure-hosted OpenAI models, including GPT-4 and GPT-3.5. Supports API key, Entra ID bearer token, and Azure credential chain authentication. | `AZURE_OPENAI_ENDPOINT`, `AZURE_OPENAI_DEPLOYMENT_NAME`, `AZURE_OPENAI_API_KEY` (optional), `AZURE_OPENAI_AD_TOKEN` (optional) | | [ChatGPT Codex](https://chatgpt.com/codex) | Access GPT-5 Codex models optimized for code generation and understanding. **Requires a ChatGPT Plus/Pro subscription.** | No manual key. Uses browser-based OAuth authentication for both CLI and Desktop. | | [Databricks](https://www.databricks.com/) | Unified data analytics and AI platform for building and deploying models. | `DATABRICKS_HOST`, `DATABRICKS_TOKEN` | diff --git a/documentation/docs/guides/azure-foundry-provider.md b/documentation/docs/guides/azure-foundry-provider.md new file mode 100644 index 000000000..552d49953 --- /dev/null +++ b/documentation/docs/guides/azure-foundry-provider.md @@ -0,0 +1,101 @@ +--- +title: Azure AI Foundry +description: Use OpenAI, Anthropic, and partner model deployments from an Azure AI Foundry project +--- + +# Azure AI Foundry + +The `azure_foundry` provider connects goose to Azure AI Foundry deployments. It supports two endpoint types: + +| Endpoint | Inference surface | +|---|---| +| Foundry project: `https://.services.ai.azure.com/api/projects/` | Deployment discovery plus publisher-aware routing | +| Foundry resource: `https://.services.ai.azure.com` | OpenAI models through Responses, Claude through Anthropic Messages, and partner models through Chat Completions | +| MaaS/serverless: `https://..models.ai.azure.com` | Chat Completions for the model bound to the endpoint | + +For project endpoints, goose discovers deployments with `GET /deployments`. Deployment names can be customized; goose uses the returned `modelPublisher` to select the protocol and `modelName` to resolve model metadata such as the context window. + +Resource endpoints do not expose project deployment discovery. goose routes recognizable model or deployment names by family: `gpt-5*` and supported o-series models use Responses, `claude-*` uses Anthropic Messages, and other names use Chat Completions. Use a project endpoint when aliases do not identify their underlying model family. + +## Configuration + +| Variable | Required | Description | +|---|---:|---| +| `AZURE_FOUNDRY_ENDPOINT` | Yes | Full Foundry project or MaaS endpoint | +| `AZURE_FOUNDRY_API_KEY` | No | API key; omit it to use Azure CLI credentials | +| `AZURE_FOUNDRY_MODEL` | MaaS only | Model bound to the configured MaaS endpoint | +| `AZURE_FOUNDRY_AD_TOKEN` | No | Pre-acquired Microsoft Entra access token; takes precedence over the API key | +| `AZURE_FOUNDRY_API_VERSION` | No | Deployment discovery API version; project endpoints default to `v1` | + +Run `goose configure`, select **Configure Providers**, and choose **Azure AI Foundry**. You can also set the variables before starting goose: + +```sh +export AZURE_FOUNDRY_ENDPOINT="https://my-resource.services.ai.azure.com/api/projects/my-project" +export AZURE_FOUNDRY_API_KEY="" +goose session +``` + +For a MaaS endpoint: + +```sh +export AZURE_FOUNDRY_ENDPOINT="https://my-deployment.eastus.models.ai.azure.com" +export AZURE_FOUNDRY_API_KEY="" +export AZURE_FOUNDRY_MODEL="" +goose session +``` + +MaaS endpoints expose a single deployed model. `AZURE_FOUNDRY_MODEL` is required for these endpoints. + +## Authentication + +Authentication is selected in this order: + +1. `AZURE_FOUNDRY_AD_TOKEN` +2. `AZURE_FOUNDRY_API_KEY` +3. Azure CLI credentials + +When neither token nor key is configured, sign in with Azure CLI before starting goose: + +```sh +az login +``` + +Project and resource endpoints request a token for `https://ai.azure.com`. MaaS endpoints request a token for `https://ml.azure.com`. + +## Protocol routing + +For a project endpoint, goose routes each deployment using metadata returned by Azure: + +- publisher `OpenAI` with a Responses-compatible model (`gpt-5*` and the supported o-series) → `/openai/v1/responses` +- publisher `Anthropic` → `/anthropic/v1/messages` +- older OpenAI models and all other publishers → `/openai/v1/chat/completions` + +If deployment discovery is temporarily unavailable, recognizable Responses-compatible OpenAI and `claude-*` names use their native surfaces. Other names use Chat Completions. + +Resource endpoints use the same inference surfaces without deployment discovery: recognizable model-family names select the native protocol. This allows a custom agent's `model:` declaration to select deployments such as `gpt-5.6-sol` without rewriting the name through canonical model resolution. + +MaaS endpoints always use `/v1/chat/completions` and the model configured by `AZURE_FOUNDRY_MODEL`. + +## Model metadata and pricing + +The deployments API provides the deployment name and underlying `modelName`, `modelVersion`, and `modelPublisher`. goose uses the underlying model name to look up a context window in its bundled model catalog. An explicit `GOOSE_CONTEXT_LIMIT` or session override still takes precedence. + +Azure pricing depends on region, SKU, offer, deployment type, and contract. The deployments API does not provide a reliable per-token price, so this provider does not attach a price to discovered deployments. + +## Troubleshooting + +### 401 or 403 + +- Ensure the key belongs to the configured endpoint. +- For Entra authentication, run `az login` again and verify that your identity has access to the Foundry project. +- Do not use a project endpoint key with a MaaS endpoint, or the reverse. + +### No deployments are listed + +- Confirm that the endpoint includes `/api/projects/`. +- Confirm that the project contains model deployments. +- If your project uses a non-default deployment API version, set `AZURE_FOUNDRY_API_VERSION`. + +### Wrong protocol for a custom deployment name + +Refresh the provider model list so goose can retrieve `modelPublisher`. Without deployment metadata, routing can only use recognizable model-name prefixes. diff --git a/goose-self-test.yaml b/goose-self-test.yaml index cfcb042b8..4e71290d1 100644 --- a/goose-self-test.yaml +++ b/goose-self-test.yaml @@ -8,6 +8,7 @@ activities: - Initialize test workspace and logging infrastructure - Test file operations (create, read, update, delete, undo) - Validate shell command execution and error handling + - Validate Azure AI Foundry provider routing and deployment aliases - Analyze code structure and parsing capabilities - Test extension discovery and management - Test load tool for knowledge injection and discovery @@ -156,6 +157,14 @@ prompt: | 4. Test symbol focus and call graphs 5. Verify LOC, function, and class counting + ### Azure AI Foundry Provider Validation + When running from a goose source checkout: + 1. Run the goose-providers Azure Foundry test module. + 2. Verify project deployments route Responses-compatible OpenAI models to Responses, Anthropic models to Messages, and partner or older OpenAI models to Chat Completions. + 3. Verify a custom deployment alias remains the wire model while the underlying Azure model controls reasoning capabilities. + 4. Verify MaaS requires its bound model and uses /v1/chat/completions. + If the source checkout or Rust toolchain is unavailable, mark this validation as skipped rather than failed. + Log results to: {{ workspace_dir }}/phase1_basic_tools.md {% endif %} diff --git a/ui/desktop/src/components/settings/providers/modal/subcomponents/ProviderLogo.tsx b/ui/desktop/src/components/settings/providers/modal/subcomponents/ProviderLogo.tsx index 12a1e8c8c..b1bdcdc40 100644 --- a/ui/desktop/src/components/settings/providers/modal/subcomponents/ProviderLogo.tsx +++ b/ui/desktop/src/components/settings/providers/modal/subcomponents/ProviderLogo.tsx @@ -9,6 +9,7 @@ import SnowflakeLogo from './icons/snowflake@3x.png'; import XaiLogo from './icons/xai@3x.png'; import MiniMaxLogo from './icons/minimax@3x.png'; import TanzuLogo from './icons/tanzu@3x.png'; +import AzureFoundryLogo from './icons/azure_foundry@3x.png'; import DefaultLogo from './icons/default@3x.png'; import { defineMessages, useIntl } from '../../../../../i18n'; @@ -32,6 +33,7 @@ const providerLogos: Record = { xai: XaiLogo, minimax: MiniMaxLogo, tanzu_ai: TanzuLogo, + azure_foundry: AzureFoundryLogo, default: DefaultLogo, }; diff --git a/ui/desktop/src/components/settings/providers/modal/subcomponents/icons/azure_foundry.png b/ui/desktop/src/components/settings/providers/modal/subcomponents/icons/azure_foundry.png new file mode 100644 index 000000000..bb60c3d69 Binary files /dev/null and b/ui/desktop/src/components/settings/providers/modal/subcomponents/icons/azure_foundry.png differ diff --git a/ui/desktop/src/components/settings/providers/modal/subcomponents/icons/azure_foundry@2x.png b/ui/desktop/src/components/settings/providers/modal/subcomponents/icons/azure_foundry@2x.png new file mode 100644 index 000000000..77e8920ba Binary files /dev/null and b/ui/desktop/src/components/settings/providers/modal/subcomponents/icons/azure_foundry@2x.png differ diff --git a/ui/desktop/src/components/settings/providers/modal/subcomponents/icons/azure_foundry@3x.png b/ui/desktop/src/components/settings/providers/modal/subcomponents/icons/azure_foundry@3x.png new file mode 100644 index 000000000..a4216305a Binary files /dev/null and b/ui/desktop/src/components/settings/providers/modal/subcomponents/icons/azure_foundry@3x.png differ