fix: remove Option from model listing return types, propagate errors (#7074)
Signed-off-by: Adrian Cole <adrian@tetrate.io>
This commit is contained in:
@@ -691,15 +691,15 @@ pub async fn configure_provider_dialog() -> anyhow::Result<bool> {
|
|||||||
};
|
};
|
||||||
spin.stop(style("Model fetch complete").green());
|
spin.stop(style("Model fetch complete").green());
|
||||||
|
|
||||||
// Select a model: on fetch error show styled error and abort; if Some(models), show list; if None, free-text input
|
// Select a model: on fetch error show styled error and abort; if models available, show list; otherwise free-text input
|
||||||
let model: String = match models_res {
|
let model: String = match models_res {
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
// Provider hook error
|
// Provider hook error
|
||||||
cliclack::outro(style(e.to_string()).on_red().white())?;
|
cliclack::outro(style(e.to_string()).on_red().white())?;
|
||||||
return Ok(false);
|
return Ok(false);
|
||||||
}
|
}
|
||||||
Ok(Some(models)) => select_model_from_list(&models, provider_meta)?,
|
Ok(models) if !models.is_empty() => select_model_from_list(&models, provider_meta)?,
|
||||||
Ok(None) => {
|
Ok(_) => {
|
||||||
let default_model =
|
let default_model =
|
||||||
std::env::var("GOOSE_MODEL").unwrap_or(provider_meta.default_model.clone());
|
std::env::var("GOOSE_MODEL").unwrap_or(provider_meta.default_model.clone());
|
||||||
cliclack::input("Enter a model from that provider:")
|
cliclack::input("Enter a model from that provider:")
|
||||||
|
|||||||
@@ -372,19 +372,6 @@ pub async fn providers() -> Result<Json<Vec<ProviderDetails>>, ErrorResponse> {
|
|||||||
pub async fn get_provider_models(
|
pub async fn get_provider_models(
|
||||||
Path(name): Path<String>,
|
Path(name): Path<String>,
|
||||||
) -> Result<Json<Vec<String>>, ErrorResponse> {
|
) -> Result<Json<Vec<String>>, ErrorResponse> {
|
||||||
let loaded_provider = goose::config::declarative_providers::load_provider(name.as_str()).ok();
|
|
||||||
// TODO(Douwe): support a get models url for custom providers
|
|
||||||
if let Some(loaded_provider) = loaded_provider {
|
|
||||||
return Ok(Json(
|
|
||||||
loaded_provider
|
|
||||||
.config
|
|
||||||
.models
|
|
||||||
.into_iter()
|
|
||||||
.map(|m| m.name)
|
|
||||||
.collect::<Vec<_>>(),
|
|
||||||
));
|
|
||||||
}
|
|
||||||
|
|
||||||
let all = get_providers().await.into_iter().collect::<Vec<_>>();
|
let all = get_providers().await.into_iter().collect::<Vec<_>>();
|
||||||
let Some((metadata, provider_type)) = all.into_iter().find(|(m, _)| m.name == name) else {
|
let Some((metadata, provider_type)) = all.into_iter().find(|(m, _)| m.name == name) else {
|
||||||
return Err(ErrorResponse::bad_request(format!(
|
return Err(ErrorResponse::bad_request(format!(
|
||||||
@@ -405,8 +392,7 @@ pub async fn get_provider_models(
|
|||||||
let models_result = provider.fetch_recommended_models().await;
|
let models_result = provider.fetch_recommended_models().await;
|
||||||
|
|
||||||
match models_result {
|
match models_result {
|
||||||
Ok(Some(models)) => Ok(Json(models)),
|
Ok(models) => Ok(Json(models)),
|
||||||
Ok(None) => Ok(Json(Vec::new())),
|
|
||||||
Err(provider_error) => Err(provider_error.into()),
|
Err(provider_error) => Err(provider_error.into()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -31,9 +31,12 @@ async fn enhance_model_error(error: ProviderError, provider: &Arc<dyn Provider>)
|
|||||||
return error;
|
return error;
|
||||||
}
|
}
|
||||||
|
|
||||||
let Ok(Some(models)) = provider.fetch_recommended_models().await else {
|
let Ok(models) = provider.fetch_recommended_models().await else {
|
||||||
return error;
|
return error;
|
||||||
};
|
};
|
||||||
|
if models.is_empty() {
|
||||||
|
return error;
|
||||||
|
}
|
||||||
|
|
||||||
ProviderError::RequestFailed(format!(
|
ProviderError::RequestFailed(format!(
|
||||||
"{}. Available models for this provider: {}",
|
"{}. Available models for this provider: {}",
|
||||||
|
|||||||
@@ -256,7 +256,7 @@ impl Provider for AnthropicProvider {
|
|||||||
Ok((message, provider_usage))
|
Ok((message, provider_usage))
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let response = self.api_client.request(None, "v1/models").api_get().await?;
|
let response = self.api_client.request(None, "v1/models").api_get().await?;
|
||||||
|
|
||||||
if response.status != StatusCode::OK {
|
if response.status != StatusCode::OK {
|
||||||
@@ -267,17 +267,18 @@ impl Provider for AnthropicProvider {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let json = response.payload.unwrap_or_default();
|
let json = response.payload.unwrap_or_default();
|
||||||
let arr = match json.get("data").and_then(|v| v.as_array()) {
|
let arr = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| {
|
||||||
Some(arr) => arr,
|
ProviderError::RequestFailed(
|
||||||
None => return Ok(None),
|
"Missing 'data' array in Anthropic models response".to_string(),
|
||||||
};
|
)
|
||||||
|
})?;
|
||||||
|
|
||||||
let mut models: Vec<String> = arr
|
let mut models: Vec<String> = arr
|
||||||
.iter()
|
.iter()
|
||||||
.filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(str::to_string))
|
.filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(str::to_string))
|
||||||
.collect();
|
.collect();
|
||||||
models.sort();
|
models.sort();
|
||||||
Ok(Some(models))
|
Ok(models)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn stream(
|
async fn stream(
|
||||||
|
|||||||
@@ -31,7 +31,9 @@ pub async fn detect_provider_from_api_key(api_key: &str) -> Option<(String, Vec<
|
|||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
Ok(Some(models)) => Some((provider_name.to_string(), models)),
|
Ok(models) if !models.is_empty() => {
|
||||||
|
Some((provider_name.to_string(), models))
|
||||||
|
}
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -447,16 +447,13 @@ pub trait Provider: Send + Sync {
|
|||||||
RetryConfig::default()
|
RetryConfig::default()
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
Ok(None)
|
Ok(vec![])
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Fetch models filtered by canonical registry and usability
|
/// Fetch models filtered by canonical registry and usability
|
||||||
async fn fetch_recommended_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_recommended_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let all_models = match self.fetch_supported_models().await? {
|
let all_models = self.fetch_supported_models().await?;
|
||||||
Some(models) => models,
|
|
||||||
None => return Ok(None),
|
|
||||||
};
|
|
||||||
|
|
||||||
let registry = CanonicalModelRegistry::bundled().map_err(|e| {
|
let registry = CanonicalModelRegistry::bundled().map_err(|e| {
|
||||||
ProviderError::ExecutionError(format!("Failed to load canonical registry: {}", e))
|
ProviderError::ExecutionError(format!("Failed to load canonical registry: {}", e))
|
||||||
@@ -501,9 +498,9 @@ pub trait Provider: Send + Sync {
|
|||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
if recommended_models.is_empty() {
|
if recommended_models.is_empty() {
|
||||||
Ok(Some(all_models))
|
Ok(all_models)
|
||||||
} else {
|
} else {
|
||||||
Ok(Some(recommended_models))
|
Ok(recommended_models)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -301,6 +301,10 @@ impl Provider for BedrockProvider {
|
|||||||
self.model.clone()
|
self.model.clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
|
Ok(BEDROCK_KNOWN_MODELS.iter().map(|s| s.to_string()).collect())
|
||||||
|
}
|
||||||
|
|
||||||
#[tracing::instrument(
|
#[tracing::instrument(
|
||||||
skip(self, model_config, system, messages, tools),
|
skip(self, model_config, system, messages, tools),
|
||||||
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
||||||
|
|||||||
@@ -493,14 +493,10 @@ async fn check_provider(
|
|||||||
};
|
};
|
||||||
|
|
||||||
let fetched_models = match provider.fetch_supported_models().await {
|
let fetched_models = match provider.fetch_supported_models().await {
|
||||||
Ok(Some(models)) => {
|
Ok(models) => {
|
||||||
println!(" ✓ Fetched {} models", models.len());
|
println!(" ✓ Fetched {} models", models.len());
|
||||||
models
|
models
|
||||||
}
|
}
|
||||||
Ok(None) => {
|
|
||||||
println!(" ⚠ Provider does not support model listing");
|
|
||||||
Vec::new()
|
|
||||||
}
|
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
println!(" ⚠ Failed to fetch models: {}", e);
|
println!(" ⚠ Failed to fetch models: {}", e);
|
||||||
println!(" This is expected if credentials are not configured.");
|
println!(" This is expected if credentials are not configured.");
|
||||||
@@ -509,11 +505,10 @@ async fn check_provider(
|
|||||||
};
|
};
|
||||||
|
|
||||||
let recommended_models = match provider.fetch_recommended_models().await {
|
let recommended_models = match provider.fetch_recommended_models().await {
|
||||||
Ok(Some(models)) => {
|
Ok(models) => {
|
||||||
println!(" ✓ Found {} recommended models", models.len());
|
println!(" ✓ Found {} recommended models", models.len());
|
||||||
models
|
models
|
||||||
}
|
}
|
||||||
Ok(None) => Vec::new(),
|
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
println!(" ⚠ Failed to fetch recommended models: {}", e);
|
println!(" ⚠ Failed to fetch recommended models: {}", e);
|
||||||
Vec::new()
|
Vec::new()
|
||||||
|
|||||||
@@ -985,13 +985,11 @@ impl Provider for ChatGptCodexProvider {
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
Ok(Some(
|
Ok(CHATGPT_CODEX_KNOWN_MODELS
|
||||||
CHATGPT_CODEX_KNOWN_MODELS
|
.iter()
|
||||||
.iter()
|
.map(|s| s.to_string())
|
||||||
.map(|s| s.to_string())
|
.collect())
|
||||||
.collect(),
|
|
||||||
))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -493,6 +493,13 @@ impl Provider for ClaudeCodeProvider {
|
|||||||
self.model.clone()
|
self.model.clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
|
Ok(CLAUDE_CODE_KNOWN_MODELS
|
||||||
|
.iter()
|
||||||
|
.map(|s| s.to_string())
|
||||||
|
.collect())
|
||||||
|
}
|
||||||
|
|
||||||
#[tracing::instrument(
|
#[tracing::instrument(
|
||||||
skip(self, model_config, system, messages, tools),
|
skip(self, model_config, system, messages, tools),
|
||||||
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
||||||
|
|||||||
@@ -662,10 +662,8 @@ impl Provider for CodexProvider {
|
|||||||
))
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
Ok(Some(
|
Ok(CODEX_KNOWN_MODELS.iter().map(|s| s.to_string()).collect())
|
||||||
CODEX_KNOWN_MODELS.iter().map(|s| s.to_string()).collect(),
|
|
||||||
))
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -355,6 +355,13 @@ impl Provider for CursorAgentProvider {
|
|||||||
self.model.clone()
|
self.model.clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
|
Ok(CURSOR_AGENT_KNOWN_MODELS
|
||||||
|
.iter()
|
||||||
|
.map(|s| s.to_string())
|
||||||
|
.collect())
|
||||||
|
}
|
||||||
|
|
||||||
#[tracing::instrument(
|
#[tracing::instrument(
|
||||||
skip(self, model_config, system, messages, tools),
|
skip(self, model_config, system, messages, tools),
|
||||||
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
||||||
|
|||||||
@@ -391,51 +391,38 @@ impl Provider for DatabricksProvider {
|
|||||||
.map_err(|e| ProviderError::ExecutionError(e.to_string()))
|
.map_err(|e| ProviderError::ExecutionError(e.to_string()))
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let response = match self
|
let response = self
|
||||||
.api_client
|
.api_client
|
||||||
.request(None, "api/2.0/serving-endpoints")
|
.request(None, "api/2.0/serving-endpoints")
|
||||||
.response_get()
|
.response_get()
|
||||||
.await
|
.await
|
||||||
{
|
.map_err(|e| {
|
||||||
Ok(resp) => resp,
|
ProviderError::RequestFailed(format!("Failed to fetch Databricks models: {}", e))
|
||||||
Err(e) => {
|
})?;
|
||||||
tracing::warn!("Failed to fetch Databricks models: {}", e);
|
|
||||||
return Ok(None);
|
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
if !response.status().is_success() {
|
if !response.status().is_success() {
|
||||||
let status = response.status();
|
let status = response.status();
|
||||||
if let Ok(error_text) = response.text().await {
|
let detail = response.text().await.unwrap_or_default();
|
||||||
tracing::warn!(
|
return Err(ProviderError::RequestFailed(format!(
|
||||||
"Failed to fetch Databricks models: {} - {}",
|
"Failed to fetch Databricks models: {} {}",
|
||||||
status,
|
status, detail
|
||||||
error_text
|
)));
|
||||||
);
|
|
||||||
} else {
|
|
||||||
tracing::warn!("Failed to fetch Databricks models: {}", status);
|
|
||||||
}
|
|
||||||
return Ok(None);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let json: Value = match response.json().await {
|
let json: Value = response.json().await.map_err(|e| {
|
||||||
Ok(json) => json,
|
ProviderError::RequestFailed(format!("Failed to parse Databricks API response: {}", e))
|
||||||
Err(e) => {
|
})?;
|
||||||
tracing::warn!("Failed to parse Databricks API response: {}", e);
|
|
||||||
return Ok(None);
|
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
let endpoints = match json.get("endpoints").and_then(|v| v.as_array()) {
|
let endpoints = json
|
||||||
Some(endpoints) => endpoints,
|
.get("endpoints")
|
||||||
None => {
|
.and_then(|v| v.as_array())
|
||||||
tracing::warn!(
|
.ok_or_else(|| {
|
||||||
|
ProviderError::RequestFailed(
|
||||||
"Unexpected response format from Databricks API: missing 'endpoints' array"
|
"Unexpected response format from Databricks API: missing 'endpoints' array"
|
||||||
);
|
.to_string(),
|
||||||
return Ok(None);
|
)
|
||||||
}
|
})?;
|
||||||
};
|
|
||||||
|
|
||||||
let models: Vec<String> = endpoints
|
let models: Vec<String> = endpoints
|
||||||
.iter()
|
.iter()
|
||||||
@@ -447,11 +434,7 @@ impl Provider for DatabricksProvider {
|
|||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
if models.is_empty() {
|
Ok(models)
|
||||||
Ok(None)
|
|
||||||
} else {
|
|
||||||
Ok(Some(models))
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -695,10 +695,10 @@ impl Provider for GcpVertexAIProvider {
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let models: Vec<String> = KNOWN_MODELS.iter().map(|s| s.to_string()).collect();
|
let models: Vec<String> = KNOWN_MODELS.iter().map(|s| s.to_string()).collect();
|
||||||
let filtered = self.filter_by_org_policy(models).await;
|
let filtered = self.filter_by_org_policy(models).await;
|
||||||
Ok(Some(filtered))
|
Ok(filtered)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -267,6 +267,13 @@ impl Provider for GeminiCliProvider {
|
|||||||
self.model.clone()
|
self.model.clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
|
Ok(GEMINI_CLI_KNOWN_MODELS
|
||||||
|
.iter()
|
||||||
|
.map(|s| s.to_string())
|
||||||
|
.collect())
|
||||||
|
}
|
||||||
|
|
||||||
#[tracing::instrument(
|
#[tracing::instrument(
|
||||||
skip(self, _model_config, system, messages, tools),
|
skip(self, _model_config, system, messages, tools),
|
||||||
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
||||||
|
|||||||
@@ -495,7 +495,7 @@ impl Provider for GithubCopilotProvider {
|
|||||||
stream_openai_compat(response, log)
|
stream_openai_compat(response, log)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let (endpoint, token) = self.get_api_info().await?;
|
let (endpoint, token) = self.get_api_info().await?;
|
||||||
let url = format!("{}/models", endpoint);
|
let url = format!("{}/models", endpoint);
|
||||||
|
|
||||||
@@ -515,10 +515,11 @@ impl Provider for GithubCopilotProvider {
|
|||||||
|
|
||||||
let json: serde_json::Value = response.json().await?;
|
let json: serde_json::Value = response.json().await?;
|
||||||
|
|
||||||
let arr = match json.get("data").and_then(|v| v.as_array()) {
|
let arr = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| {
|
||||||
Some(arr) => arr,
|
ProviderError::RequestFailed(
|
||||||
None => return Ok(None),
|
"Missing 'data' array in GitHub Copilot models response".to_string(),
|
||||||
};
|
)
|
||||||
|
})?;
|
||||||
let mut models: Vec<String> = arr
|
let mut models: Vec<String> = arr
|
||||||
.iter()
|
.iter()
|
||||||
.filter_map(|m| {
|
.filter_map(|m| {
|
||||||
@@ -532,7 +533,7 @@ impl Provider for GithubCopilotProvider {
|
|||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
models.sort();
|
models.sort();
|
||||||
Ok(Some(models))
|
Ok(models)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn configure_oauth(&self) -> Result<(), ProviderError> {
|
async fn configure_oauth(&self) -> Result<(), ProviderError> {
|
||||||
|
|||||||
@@ -187,24 +187,28 @@ impl Provider for GoogleProvider {
|
|||||||
Ok((message, provider_usage))
|
Ok((message, provider_usage))
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let response = self
|
let response = self
|
||||||
.api_client
|
.api_client
|
||||||
.request(None, "v1beta/models")
|
.request(None, "v1beta/models")
|
||||||
.response_get()
|
.response_get()
|
||||||
.await?;
|
.await?;
|
||||||
let json: serde_json::Value = response.json().await?;
|
let json: serde_json::Value = response.json().await?;
|
||||||
let arr = match json.get("models").and_then(|v| v.as_array()) {
|
let arr = json
|
||||||
Some(arr) => arr,
|
.get("models")
|
||||||
None => return Ok(None),
|
.and_then(|v| v.as_array())
|
||||||
};
|
.ok_or_else(|| {
|
||||||
|
ProviderError::RequestFailed(
|
||||||
|
"Missing 'models' array in Google models response".to_string(),
|
||||||
|
)
|
||||||
|
})?;
|
||||||
let mut models: Vec<String> = arr
|
let mut models: Vec<String> = arr
|
||||||
.iter()
|
.iter()
|
||||||
.filter_map(|m| m.get("name").and_then(|v| v.as_str()))
|
.filter_map(|m| m.get("name").and_then(|v| v.as_str()))
|
||||||
.map(|name| name.split('/').next_back().unwrap_or(name).to_string())
|
.map(|name| name.split('/').next_back().unwrap_or(name).to_string())
|
||||||
.collect();
|
.collect();
|
||||||
models.sort();
|
models.sort();
|
||||||
Ok(Some(models))
|
Ok(models)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn supports_streaming(&self) -> bool {
|
fn supports_streaming(&self) -> bool {
|
||||||
|
|||||||
@@ -445,22 +445,14 @@ impl Provider for LeadWorkerProvider {
|
|||||||
final_result
|
final_result
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
// Combine models from both providers
|
// Combine models from both providers
|
||||||
let lead_models = self.lead_provider.fetch_supported_models().await?;
|
let mut all_models = self.lead_provider.fetch_supported_models().await?;
|
||||||
let worker_models = self.worker_provider.fetch_supported_models().await?;
|
let worker_models = self.worker_provider.fetch_supported_models().await?;
|
||||||
|
all_models.extend(worker_models);
|
||||||
match (lead_models, worker_models) {
|
all_models.sort();
|
||||||
(Some(lead), Some(worker)) => {
|
all_models.dedup();
|
||||||
let mut all_models = lead;
|
Ok(all_models)
|
||||||
all_models.extend(worker);
|
|
||||||
all_models.sort();
|
|
||||||
all_models.dedup();
|
|
||||||
Ok(Some(all_models))
|
|
||||||
}
|
|
||||||
(Some(models), None) | (None, Some(models)) => Ok(Some(models)),
|
|
||||||
(None, None) => Ok(None),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn supports_embeddings(&self) -> bool {
|
fn supports_embeddings(&self) -> bool {
|
||||||
|
|||||||
@@ -223,17 +223,11 @@ impl Provider for LiteLLMProvider {
|
|||||||
self.model.model_name.to_lowercase().contains("claude")
|
self.model.model_name.to_lowercase().contains("claude")
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
match self.fetch_models().await {
|
let models = self.fetch_models().await.map_err(|e| {
|
||||||
Ok(models) => {
|
ProviderError::RequestFailed(format!("Failed to fetch models from LiteLLM: {}", e))
|
||||||
let model_names: Vec<String> = models.into_iter().map(|m| m.name).collect();
|
})?;
|
||||||
Ok(Some(model_names))
|
Ok(models.into_iter().map(|m| m.name).collect())
|
||||||
}
|
|
||||||
Err(e) => {
|
|
||||||
tracing::warn!("Failed to fetch models from LiteLLM: {}", e);
|
|
||||||
Ok(None)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -307,7 +307,7 @@ impl Provider for OllamaProvider {
|
|||||||
stream_ollama(response, log)
|
stream_ollama(response, log)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let response = self
|
let response = self
|
||||||
.api_client
|
.api_client
|
||||||
.request(None, "api/tags")
|
.request(None, "api/tags")
|
||||||
@@ -340,7 +340,7 @@ impl Provider for OllamaProvider {
|
|||||||
|
|
||||||
model_names.sort();
|
model_names.sort();
|
||||||
|
|
||||||
Ok(Some(model_names))
|
Ok(model_names)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -353,7 +353,7 @@ impl Provider for OpenAiProvider {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let models_path = self.base_path.replace("v1/chat/completions", "v1/models");
|
let models_path = self.base_path.replace("v1/chat/completions", "v1/models");
|
||||||
let response = self
|
let response = self
|
||||||
.api_client
|
.api_client
|
||||||
@@ -377,7 +377,7 @@ impl Provider for OpenAiProvider {
|
|||||||
.filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(str::to_string))
|
.filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(str::to_string))
|
||||||
.collect();
|
.collect();
|
||||||
models.sort();
|
models.sort();
|
||||||
Ok(Some(models))
|
Ok(models)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn supports_embeddings(&self) -> bool {
|
fn supports_embeddings(&self) -> bool {
|
||||||
|
|||||||
@@ -112,7 +112,7 @@ impl Provider for OpenAiCompatibleProvider {
|
|||||||
Ok((message, ProviderUsage::new(response_model, usage)))
|
Ok((message, ProviderUsage::new(response_model, usage)))
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let response = self
|
let response = self
|
||||||
.api_client
|
.api_client
|
||||||
.response_get(None, "models")
|
.response_get(None, "models")
|
||||||
@@ -128,18 +128,15 @@ impl Provider for OpenAiCompatibleProvider {
|
|||||||
return Err(ProviderError::Authentication(msg.to_string()));
|
return Err(ProviderError::Authentication(msg.to_string()));
|
||||||
}
|
}
|
||||||
|
|
||||||
let data = json.get("data").and_then(|v| v.as_array());
|
let arr = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| {
|
||||||
match data {
|
ProviderError::RequestFailed("Missing 'data' array in models response".to_string())
|
||||||
Some(arr) => {
|
})?;
|
||||||
let mut models: Vec<String> = arr
|
let mut models: Vec<String> = arr
|
||||||
.iter()
|
.iter()
|
||||||
.filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(str::to_string))
|
.filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(str::to_string))
|
||||||
.collect();
|
.collect();
|
||||||
models.sort();
|
models.sort();
|
||||||
Ok(Some(models))
|
Ok(models)
|
||||||
}
|
|
||||||
None => Ok(None),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn supports_streaming(&self) -> bool {
|
fn supports_streaming(&self) -> bool {
|
||||||
|
|||||||
@@ -309,39 +309,35 @@ impl Provider for OpenRouterProvider {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Fetch supported models from OpenRouter API (only models with tool support)
|
/// Fetch supported models from OpenRouter API (only models with tool support)
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
// Handle request failures gracefully
|
let response = self
|
||||||
// If the request fails, fall back to manual entry
|
|
||||||
let response = match self
|
|
||||||
.api_client
|
.api_client
|
||||||
.request(None, "api/v1/models")
|
.request(None, "api/v1/models")
|
||||||
.response_get()
|
.response_get()
|
||||||
.await
|
.await
|
||||||
{
|
.map_err(|e| {
|
||||||
Ok(response) => response,
|
ProviderError::RequestFailed(format!(
|
||||||
Err(e) => {
|
"Failed to fetch models from OpenRouter API: {}",
|
||||||
tracing::warn!("Failed to fetch models from OpenRouter API: {}, falling back to manual model entry", e);
|
e
|
||||||
return Ok(None);
|
))
|
||||||
}
|
})?;
|
||||||
};
|
|
||||||
|
|
||||||
// Handle JSON parsing failures gracefully
|
let json: serde_json::Value = response.json().await.map_err(|e| {
|
||||||
let json: serde_json::Value = match response.json().await {
|
ProviderError::RequestFailed(format!(
|
||||||
Ok(json) => json,
|
"Failed to parse OpenRouter API response as JSON: {}",
|
||||||
Err(e) => {
|
e
|
||||||
tracing::warn!("Failed to parse OpenRouter API response as JSON: {}, falling back to manual model entry", e);
|
))
|
||||||
return Ok(None);
|
})?;
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
// Check for error in response
|
|
||||||
if let Some(err_obj) = json.get("error") {
|
if let Some(err_obj) = json.get("error") {
|
||||||
let msg = err_obj
|
let msg = err_obj
|
||||||
.get("message")
|
.get("message")
|
||||||
.and_then(|v| v.as_str())
|
.and_then(|v| v.as_str())
|
||||||
.unwrap_or("unknown error");
|
.unwrap_or("unknown error");
|
||||||
tracing::warn!("OpenRouter API returned an error: {}", msg);
|
return Err(ProviderError::RequestFailed(format!(
|
||||||
return Ok(None);
|
"OpenRouter API returned an error: {}",
|
||||||
|
msg
|
||||||
|
)));
|
||||||
}
|
}
|
||||||
|
|
||||||
let data = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| {
|
let data = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| {
|
||||||
@@ -380,14 +376,8 @@ impl Provider for OpenRouterProvider {
|
|||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
// If no models with tool support were found, fall back to manual entry
|
|
||||||
if models.is_empty() {
|
|
||||||
tracing::warn!("No models with tool support found in OpenRouter API response, falling back to manual model entry");
|
|
||||||
return Ok(None);
|
|
||||||
}
|
|
||||||
|
|
||||||
models.sort();
|
models.sort();
|
||||||
Ok(Some(models))
|
Ok(models)
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn supports_cache_control(&self) -> bool {
|
async fn supports_cache_control(&self) -> bool {
|
||||||
|
|||||||
@@ -328,6 +328,13 @@ impl Provider for SnowflakeProvider {
|
|||||||
self.model.clone()
|
self.model.clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
|
Ok(SNOWFLAKE_KNOWN_MODELS
|
||||||
|
.iter()
|
||||||
|
.map(|s| s.to_string())
|
||||||
|
.collect())
|
||||||
|
}
|
||||||
|
|
||||||
#[tracing::instrument(
|
#[tracing::instrument(
|
||||||
skip(self, model_config, system, messages, tools),
|
skip(self, model_config, system, messages, tools),
|
||||||
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
|
||||||
|
|||||||
@@ -244,7 +244,7 @@ impl Provider for TetrateProvider {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Fetch supported models from Tetrate Agent Router Service API (only models with tool support)
|
/// Fetch supported models from Tetrate Agent Router Service API (only models with tool support)
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
// Use the existing api_client which already has authentication configured
|
// Use the existing api_client which already has authentication configured
|
||||||
let response = match self
|
let response = match self
|
||||||
.api_client
|
.api_client
|
||||||
@@ -261,14 +261,12 @@ impl Provider for TetrateProvider {
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
// Handle JSON parsing failures gracefully
|
let json: serde_json::Value = response.json().await.map_err(|e| {
|
||||||
let json: serde_json::Value = match response.json().await {
|
ProviderError::ExecutionError(format!(
|
||||||
Ok(json) => json,
|
"Failed to parse Tetrate API response: {}. Please check your API key and account at {}",
|
||||||
Err(e) => {
|
e, TETRATE_DOC_URL
|
||||||
tracing::warn!("Failed to parse Tetrate Agent Router Service API response as JSON: {}, falling back to manual model entry", e);
|
))
|
||||||
return Ok(None);
|
})?;
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
// Check for error in response
|
// Check for error in response
|
||||||
if let Some(err_obj) = json.get("error") {
|
if let Some(err_obj) = json.get("error") {
|
||||||
@@ -276,10 +274,6 @@ impl Provider for TetrateProvider {
|
|||||||
.get("message")
|
.get("message")
|
||||||
.and_then(|v| v.as_str())
|
.and_then(|v| v.as_str())
|
||||||
.unwrap_or("unknown error");
|
.unwrap_or("unknown error");
|
||||||
tracing::warn!(
|
|
||||||
"Tetrate Agent Router Service API returned an error: {}",
|
|
||||||
msg
|
|
||||||
);
|
|
||||||
return Err(ProviderError::ExecutionError(format!(
|
return Err(ProviderError::ExecutionError(format!(
|
||||||
"Tetrate API error: {}. Please check your API key and account at {}",
|
"Tetrate API error: {}. Please check your API key and account at {}",
|
||||||
msg, TETRATE_DOC_URL
|
msg, TETRATE_DOC_URL
|
||||||
@@ -288,13 +282,12 @@ impl Provider for TetrateProvider {
|
|||||||
|
|
||||||
// The response format from /v1/models is expected to be OpenAI-compatible
|
// The response format from /v1/models is expected to be OpenAI-compatible
|
||||||
// It should have a "data" field with an array of model objects
|
// It should have a "data" field with an array of model objects
|
||||||
let data = match json.get("data").and_then(|v| v.as_array()) {
|
let data = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| {
|
||||||
Some(data) => data,
|
ProviderError::ExecutionError(format!(
|
||||||
None => {
|
"Tetrate API response missing 'data' field. Please check your API key and account at {}",
|
||||||
tracing::warn!("Tetrate Agent Router Service API response missing 'data' field, falling back to manual model entry");
|
TETRATE_DOC_URL
|
||||||
return Ok(None);
|
))
|
||||||
}
|
})?;
|
||||||
};
|
|
||||||
|
|
||||||
let mut models: Vec<String> = data
|
let mut models: Vec<String> = data
|
||||||
.iter()
|
.iter()
|
||||||
@@ -312,13 +305,8 @@ impl Provider for TetrateProvider {
|
|||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
if models.is_empty() {
|
|
||||||
tracing::warn!("No models found in Tetrate Agent Router Service API response, falling back to manual model entry");
|
|
||||||
return Ok(None);
|
|
||||||
}
|
|
||||||
|
|
||||||
models.sort();
|
models.sort();
|
||||||
Ok(Some(models))
|
Ok(models)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn supports_streaming(&self) -> bool {
|
fn supports_streaming(&self) -> bool {
|
||||||
|
|||||||
@@ -239,7 +239,7 @@ impl Provider for VeniceProvider {
|
|||||||
self.model.clone()
|
self.model.clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
|
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||||
let response = self
|
let response = self
|
||||||
.api_client
|
.api_client
|
||||||
.request(None, &self.models_path)
|
.request(None, &self.models_path)
|
||||||
@@ -264,7 +264,7 @@ impl Provider for VeniceProvider {
|
|||||||
})
|
})
|
||||||
.collect::<Vec<String>>();
|
.collect::<Vec<String>>();
|
||||||
models.sort();
|
models.sort();
|
||||||
Ok(Some(models))
|
Ok(models)
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tracing::instrument(
|
#[tracing::instrument(
|
||||||
|
|||||||
@@ -374,19 +374,17 @@ impl ProviderTester {
|
|||||||
dbg!(&models);
|
dbg!(&models);
|
||||||
println!("===================");
|
println!("===================");
|
||||||
|
|
||||||
if let Some(models) = models {
|
assert!(!models.is_empty(), "Expected non-empty model list");
|
||||||
assert!(!models.is_empty(), "Expected non-empty model list");
|
let model_name = &self.provider.get_model_config().model_name;
|
||||||
let model_name = &self.provider.get_model_config().model_name;
|
// Some providers (e.g. Ollama) return names with tags like "qwen3:latest"
|
||||||
// Some providers (e.g. Ollama) return names with tags like "qwen3:latest"
|
// while the configured model name may be just "qwen3".
|
||||||
// while the configured model name may be just "qwen3".
|
assert!(
|
||||||
assert!(
|
models
|
||||||
models
|
.iter()
|
||||||
.iter()
|
.any(|m| m == model_name || m.starts_with(&format!("{}:", model_name))),
|
||||||
.any(|m| m == model_name || m.starts_with(&format!("{}:", model_name))),
|
"Expected model '{}' in supported models",
|
||||||
"Expected model '{}' in supported models",
|
model_name
|
||||||
model_name
|
);
|
||||||
);
|
|
||||||
}
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user