Better search paths and handling of CLI providers (#5554)
Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
This commit is contained in:
@@ -344,6 +344,8 @@ derive_utoipa!(Icon as IconSchema);
|
||||
super::routes::config_management::get_custom_provider,
|
||||
super::routes::config_management::update_custom_provider,
|
||||
super::routes::config_management::remove_custom_provider,
|
||||
super::routes::config_management::check_provider,
|
||||
super::routes::config_management::set_config_provider,
|
||||
super::routes::agent::start_agent,
|
||||
super::routes::agent::resume_agent,
|
||||
super::routes::agent::get_tools,
|
||||
@@ -394,6 +396,8 @@ derive_utoipa!(Icon as IconSchema);
|
||||
super::routes::config_management::ToolPermission,
|
||||
super::routes::config_management::UpsertPermissionsQuery,
|
||||
super::routes::config_management::UpdateCustomProviderRequest,
|
||||
super::routes::config_management::CheckProviderRequest,
|
||||
super::routes::config_management::SetProviderRequest,
|
||||
super::routes::reply::PermissionConfirmationRequest,
|
||||
super::routes::reply::ChatRequest,
|
||||
super::routes::session::ImportSessionRequest,
|
||||
|
||||
@@ -3,6 +3,7 @@ use crate::routes::recipe_utils::{
|
||||
apply_recipe_to_agent, build_recipe_with_parameter_values, load_recipe_by_id, validate_recipe,
|
||||
};
|
||||
use crate::state::AppState;
|
||||
use axum::response::IntoResponse;
|
||||
use axum::{
|
||||
extract::{Query, State},
|
||||
http::StatusCode,
|
||||
@@ -399,36 +400,42 @@ async fn get_tools(
|
||||
async fn update_agent_provider(
|
||||
State(state): State<Arc<AppState>>,
|
||||
Json(payload): Json<UpdateProviderRequest>,
|
||||
) -> Result<StatusCode, StatusCode> {
|
||||
) -> Result<(), impl IntoResponse> {
|
||||
let agent = state
|
||||
.get_agent_for_route(payload.session_id.clone())
|
||||
.await?;
|
||||
.await
|
||||
.map_err(|e| (e, "No agent for session id".to_owned()))?;
|
||||
|
||||
let config = Config::global();
|
||||
let model = match payload.model.or_else(|| config.get_goose_model().ok()) {
|
||||
Some(m) => m,
|
||||
None => {
|
||||
tracing::error!("No model specified");
|
||||
return Err(StatusCode::BAD_REQUEST);
|
||||
return Err((StatusCode::BAD_REQUEST, "No model specified".to_owned()));
|
||||
}
|
||||
};
|
||||
|
||||
let model_config = ModelConfig::new(&model).map_err(|e| {
|
||||
tracing::error!("Invalid model config: {}", e);
|
||||
StatusCode::BAD_REQUEST
|
||||
(
|
||||
StatusCode::BAD_REQUEST,
|
||||
format!("Invalid model config: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
let new_provider = create(&payload.provider, model_config).await.map_err(|e| {
|
||||
tracing::error!("Failed to create provider: {}", e);
|
||||
StatusCode::BAD_REQUEST
|
||||
(
|
||||
StatusCode::BAD_REQUEST,
|
||||
format!("Failed to create {} provider: {}", &payload.provider, e),
|
||||
)
|
||||
})?;
|
||||
|
||||
agent.update_provider(new_provider).await.map_err(|e| {
|
||||
tracing::error!("Failed to update provider: {}", e);
|
||||
StatusCode::INTERNAL_SERVER_ERROR
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("Failed to update provider: {}", e),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(StatusCode::OK)
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
|
||||
@@ -12,6 +12,7 @@ use goose::config::ExtensionEntry;
|
||||
use goose::config::{Config, ConfigError};
|
||||
use goose::model::ModelConfig;
|
||||
use goose::providers::base::{ProviderMetadata, ProviderType};
|
||||
use goose::providers::create_with_default_model;
|
||||
use goose::providers::pricing::{
|
||||
get_all_pricing, get_model_pricing, parse_model_id, refresh_pricing,
|
||||
};
|
||||
@@ -88,6 +89,17 @@ pub struct UpdateCustomProviderRequest {
|
||||
pub supports_streaming: Option<bool>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, ToSchema)]
|
||||
pub struct CheckProviderRequest {
|
||||
pub provider: String,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, ToSchema)]
|
||||
pub struct SetProviderRequest {
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
}
|
||||
|
||||
#[derive(Serialize, ToSchema)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct MaskedSecret {
|
||||
@@ -734,6 +746,41 @@ pub async fn update_custom_provider(
|
||||
Ok(Json(format!("Updated custom provider: {}", id)))
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/config/check_provider",
|
||||
request_body = CheckProviderRequest,
|
||||
)]
|
||||
pub async fn check_provider(
|
||||
Json(CheckProviderRequest { provider }): Json<CheckProviderRequest>,
|
||||
) -> Result<(), (StatusCode, String)> {
|
||||
create_with_default_model(&provider)
|
||||
.await
|
||||
.map_err(|err| (StatusCode::BAD_REQUEST, err.to_string()))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/config/set_provider",
|
||||
request_body = SetProviderRequest,
|
||||
)]
|
||||
pub async fn set_config_provider(
|
||||
Json(SetProviderRequest { provider, model }): Json<SetProviderRequest>,
|
||||
) -> Result<(), (StatusCode, String)> {
|
||||
create_with_default_model(&provider)
|
||||
.await
|
||||
.and_then(|_| {
|
||||
let config = Config::global();
|
||||
config
|
||||
.set_goose_provider(provider)
|
||||
.and_then(|_| config.set_goose_model(model))
|
||||
.map_err(|e| anyhow::anyhow!(e))
|
||||
})
|
||||
.map_err(|err| (StatusCode::BAD_REQUEST, err.to_string()))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn routes(state: Arc<AppState>) -> Router {
|
||||
Router::new()
|
||||
.route("/config", get(read_all_config))
|
||||
@@ -758,6 +805,8 @@ pub fn routes(state: Arc<AppState>) -> Router {
|
||||
)
|
||||
.route("/config/custom-providers/{id}", put(update_custom_provider))
|
||||
.route("/config/custom-providers/{id}", get(get_custom_provider))
|
||||
.route("/config/check_provider", post(check_provider))
|
||||
.route("/config/set_provider", post(set_config_provider))
|
||||
.with_state(state)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user