Expose raw provider supported models over ACP (#9475)
Signed-off-by: Bradley Axen <baxen@squareup.com> Signed-off-by: Matt Toohey <contact@matttoohey.com> Co-authored-by: Matt Toohey <contact@matttoohey.com>
This commit is contained in:
@@ -1122,6 +1122,24 @@ pub struct ListProvidersResponse {
|
||||
pub entries: Vec<ProviderInventoryEntryDto>,
|
||||
}
|
||||
|
||||
/// List the raw model identifiers returned by a provider's live supported-models API.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
|
||||
#[request(
|
||||
method = "_goose/unstable/providers/supported-models/list",
|
||||
response = ProviderSupportedModelsListResponse
|
||||
)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ProviderSupportedModelsListRequest {
|
||||
pub provider_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ProviderSupportedModelsListResponse {
|
||||
pub provider_id: String,
|
||||
pub models: Vec<String>,
|
||||
}
|
||||
|
||||
/// Trigger a background refresh of provider inventories.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
|
||||
#[request(
|
||||
|
||||
@@ -70,6 +70,11 @@
|
||||
"requestType": "ListProvidersRequest_unstable",
|
||||
"responseType": "ListProvidersResponse_unstable"
|
||||
},
|
||||
{
|
||||
"method": "_goose/unstable/providers/supported-models/list",
|
||||
"requestType": "ProviderSupportedModelsListRequest_unstable",
|
||||
"responseType": "ProviderSupportedModelsListResponse_unstable"
|
||||
},
|
||||
{
|
||||
"method": "_goose/unstable/providers/catalog/list",
|
||||
"requestType": "ProviderCatalogListRequest_unstable",
|
||||
|
||||
@@ -568,6 +568,40 @@
|
||||
],
|
||||
"description": "A single model in provider inventory."
|
||||
},
|
||||
"ProviderSupportedModelsListRequest_unstable": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"providerId": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"providerId"
|
||||
],
|
||||
"description": "List the raw model identifiers returned by a provider's live supported-models API.",
|
||||
"x-side": "agent",
|
||||
"x-method": "_goose/unstable/providers/supported-models/list"
|
||||
},
|
||||
"ProviderSupportedModelsListResponse_unstable": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"providerId": {
|
||||
"type": "string"
|
||||
},
|
||||
"models": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"providerId",
|
||||
"models"
|
||||
],
|
||||
"x-side": "agent",
|
||||
"x-method": "_goose/unstable/providers/supported-models/list"
|
||||
},
|
||||
"ProviderCatalogListRequest_unstable": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
@@ -2755,6 +2789,15 @@
|
||||
"description": "Params for _goose/unstable/providers/list",
|
||||
"title": "ListProvidersRequest_unstable"
|
||||
},
|
||||
{
|
||||
"allOf": [
|
||||
{
|
||||
"$ref": "#/$defs/ProviderSupportedModelsListRequest_unstable"
|
||||
}
|
||||
],
|
||||
"description": "Params for _goose/unstable/providers/supported-models/list",
|
||||
"title": "ProviderSupportedModelsListRequest_unstable"
|
||||
},
|
||||
{
|
||||
"allOf": [
|
||||
{
|
||||
@@ -3219,6 +3262,14 @@
|
||||
],
|
||||
"title": "ListProvidersResponse_unstable"
|
||||
},
|
||||
{
|
||||
"allOf": [
|
||||
{
|
||||
"$ref": "#/$defs/ProviderSupportedModelsListResponse_unstable"
|
||||
}
|
||||
],
|
||||
"title": "ProviderSupportedModelsListResponse_unstable"
|
||||
},
|
||||
{
|
||||
"allOf": [
|
||||
{
|
||||
|
||||
@@ -122,6 +122,14 @@ impl GooseAcpAgent {
|
||||
self.on_list_providers(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ProviderSupportedModelsListRequest)]
|
||||
async fn dispatch_list_provider_supported_models(
|
||||
&self,
|
||||
req: ProviderSupportedModelsListRequest,
|
||||
) -> Result<ProviderSupportedModelsListResponse, agent_client_protocol::Error> {
|
||||
self.on_list_provider_supported_models(req).await
|
||||
}
|
||||
|
||||
#[custom_method(ProviderCatalogListRequest)]
|
||||
async fn dispatch_list_provider_catalog(
|
||||
&self,
|
||||
|
||||
@@ -438,6 +438,30 @@ impl GooseAcpAgent {
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn on_list_provider_supported_models(
|
||||
&self,
|
||||
req: ProviderSupportedModelsListRequest,
|
||||
) -> Result<ProviderSupportedModelsListResponse, agent_client_protocol::Error> {
|
||||
let entry = crate::providers::get_from_registry(&req.provider_id)
|
||||
.await
|
||||
.invalid_params_err_ctx("Unknown provider")?;
|
||||
let model_config = crate::model::ModelConfig::new(&entry.metadata().default_model)
|
||||
.invalid_params_err_ctx("Invalid default model")?;
|
||||
let provider = self
|
||||
.create_provider(&req.provider_id, model_config, Vec::new(), None)
|
||||
.await
|
||||
.internal_err_ctx("Failed to initialize provider")?;
|
||||
let models = provider
|
||||
.fetch_supported_models()
|
||||
.await
|
||||
.internal_err_ctx("Failed to fetch provider supported models")?;
|
||||
|
||||
Ok(ProviderSupportedModelsListResponse {
|
||||
provider_id: req.provider_id,
|
||||
models,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn on_list_provider_catalog(
|
||||
&self,
|
||||
req: ProviderCatalogListRequest,
|
||||
|
||||
@@ -21,6 +21,7 @@ struct MockProvider {
|
||||
name: String,
|
||||
model_config: ModelConfig,
|
||||
recommended_models: Vec<String>,
|
||||
supported_models: Vec<String>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
@@ -47,6 +48,10 @@ impl Provider for MockProvider {
|
||||
async fn fetch_recommended_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||
Ok(self.recommended_models.clone())
|
||||
}
|
||||
|
||||
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||
Ok(self.supported_models.clone())
|
||||
}
|
||||
}
|
||||
|
||||
fn mock_provider_factory() -> AcpProviderFactory {
|
||||
@@ -62,6 +67,7 @@ fn mock_provider_factory() -> AcpProviderFactory {
|
||||
Ok(Arc::new(MockProvider {
|
||||
name: provider_name,
|
||||
model_config,
|
||||
supported_models: recommended_models.clone(),
|
||||
recommended_models,
|
||||
}) as Arc<dyn Provider>)
|
||||
})
|
||||
@@ -133,6 +139,7 @@ fn test_new_session_passes_cwd_to_provider_factory() {
|
||||
name: provider_name,
|
||||
model_config,
|
||||
recommended_models: Vec::new(),
|
||||
supported_models: Vec::new(),
|
||||
}) as Arc<dyn Provider>)
|
||||
})
|
||||
},
|
||||
@@ -180,6 +187,7 @@ fn test_load_session_passes_load_cwd_to_provider_factory() {
|
||||
name: provider_name,
|
||||
model_config,
|
||||
recommended_models: Vec::new(),
|
||||
supported_models: Vec::new(),
|
||||
}) as Arc<dyn Provider>)
|
||||
})
|
||||
},
|
||||
@@ -673,3 +681,52 @@ fn test_developer_fs_requests_use_acp_session_id() {
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_custom_provider_supported_models_lists_raw_provider_models() {
|
||||
run_test(async move {
|
||||
let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await;
|
||||
let provider_factory: AcpProviderFactory =
|
||||
Arc::new(|provider_name, model_config, _extensions, _working_dir| {
|
||||
Box::pin(async move {
|
||||
Ok(Arc::new(MockProvider {
|
||||
name: provider_name,
|
||||
model_config,
|
||||
recommended_models: vec!["canonical-filtered-model".to_string()],
|
||||
supported_models: vec![
|
||||
"goose-claude-opus-4-8".to_string(),
|
||||
"raw-databricks-endpoint".to_string(),
|
||||
],
|
||||
}) as Arc<dyn Provider>)
|
||||
})
|
||||
});
|
||||
let conn = AcpServerConnection::new(
|
||||
TestConnectionConfig {
|
||||
provider_factory: Some(provider_factory),
|
||||
..Default::default()
|
||||
},
|
||||
openai,
|
||||
)
|
||||
.await;
|
||||
|
||||
let response = send_custom(
|
||||
conn.cx(),
|
||||
"_goose/unstable/providers/supported-models/list",
|
||||
serde_json::json!({ "providerId": "openai" }),
|
||||
)
|
||||
.await
|
||||
.expect("provider supported models list should succeed");
|
||||
|
||||
assert_eq!(
|
||||
response.get("providerId"),
|
||||
Some(&serde_json::json!("openai"))
|
||||
);
|
||||
assert_eq!(
|
||||
response.get("models"),
|
||||
Some(&serde_json::json!([
|
||||
"goose-claude-opus-4-8",
|
||||
"raw-databricks-endpoint"
|
||||
]))
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user