overhaul provider inventory and agent/model selection (#8652)
Signed-off-by: Bradley Axen <baxen@squareup.com>
This commit is contained in:
@@ -51,9 +51,14 @@
|
||||
"responseType": "GetProviderDetailsResponse"
|
||||
},
|
||||
{
|
||||
"method": "_goose/providers/models",
|
||||
"requestType": "GetProviderModelsRequest",
|
||||
"responseType": "GetProviderModelsResponse"
|
||||
"method": "_goose/providers/inventory",
|
||||
"requestType": "GetProviderInventoryRequest",
|
||||
"responseType": "GetProviderInventoryResponse"
|
||||
},
|
||||
{
|
||||
"method": "_goose/providers/inventory/refresh",
|
||||
"requestType": "RefreshProviderInventoryRequest",
|
||||
"responseType": "RefreshProviderInventoryResponse"
|
||||
},
|
||||
{
|
||||
"method": "_goose/config/read",
|
||||
|
||||
@@ -362,36 +362,224 @@
|
||||
"contextLimit"
|
||||
]
|
||||
},
|
||||
"GetProviderModelsRequest": {
|
||||
"GetProviderInventoryRequest": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"providerName": {
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"providerName"
|
||||
],
|
||||
"description": "Fetch the full list of models available for a specific provider.",
|
||||
"x-side": "agent",
|
||||
"x-method": "_goose/providers/models"
|
||||
},
|
||||
"GetProviderModelsResponse": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"models": {
|
||||
"providerIds": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"description": "Only return entries for these providers. Empty means all.",
|
||||
"default": []
|
||||
}
|
||||
},
|
||||
"description": "Read per-provider inventory. Always returns immediately from stored state.",
|
||||
"x-side": "agent",
|
||||
"x-method": "_goose/providers/inventory"
|
||||
},
|
||||
"GetProviderInventoryResponse": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"entries": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/$defs/ProviderInventoryEntryDto"
|
||||
}
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"models"
|
||||
"entries"
|
||||
],
|
||||
"description": "Provider models response.",
|
||||
"description": "Provider inventory response.",
|
||||
"x-side": "agent",
|
||||
"x-method": "_goose/providers/models"
|
||||
"x-method": "_goose/providers/inventory"
|
||||
},
|
||||
"ProviderInventoryEntryDto": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"providerId": {
|
||||
"type": "string",
|
||||
"description": "Provider identifier."
|
||||
},
|
||||
"providerName": {
|
||||
"type": "string",
|
||||
"description": "Human-readable provider name."
|
||||
},
|
||||
"configured": {
|
||||
"type": "boolean",
|
||||
"description": "Whether Goose has enough configuration to use this provider."
|
||||
},
|
||||
"supportsRefresh": {
|
||||
"type": "boolean",
|
||||
"description": "Whether this provider supports background inventory refresh."
|
||||
},
|
||||
"refreshing": {
|
||||
"type": "boolean",
|
||||
"description": "Whether a refresh is currently in flight."
|
||||
},
|
||||
"models": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/$defs/ProviderInventoryModelDto"
|
||||
},
|
||||
"description": "The list of available models."
|
||||
},
|
||||
"lastUpdatedAt": {
|
||||
"type": [
|
||||
"string",
|
||||
"null"
|
||||
],
|
||||
"description": "When this entry was last successfully refreshed (ISO 8601)."
|
||||
},
|
||||
"lastRefreshAttemptAt": {
|
||||
"type": [
|
||||
"string",
|
||||
"null"
|
||||
],
|
||||
"description": "When a refresh was most recently attempted (ISO 8601)."
|
||||
},
|
||||
"lastRefreshError": {
|
||||
"type": [
|
||||
"string",
|
||||
"null"
|
||||
],
|
||||
"description": "The last refresh failure message, if any."
|
||||
},
|
||||
"stale": {
|
||||
"type": "boolean",
|
||||
"description": "Whether we believe this data may be outdated."
|
||||
},
|
||||
"modelSelectionHint": {
|
||||
"type": [
|
||||
"string",
|
||||
"null"
|
||||
],
|
||||
"description": "Guidance message shown when this provider manages its own model selection externally."
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"providerId",
|
||||
"providerName",
|
||||
"configured",
|
||||
"supportsRefresh",
|
||||
"refreshing",
|
||||
"models",
|
||||
"stale"
|
||||
],
|
||||
"description": "Provider inventory entry."
|
||||
},
|
||||
"ProviderInventoryModelDto": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"id": {
|
||||
"type": "string",
|
||||
"description": "Model identifier as the provider knows it."
|
||||
},
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "Human-readable display name."
|
||||
},
|
||||
"family": {
|
||||
"type": [
|
||||
"string",
|
||||
"null"
|
||||
],
|
||||
"description": "Model family for grouping in UI."
|
||||
},
|
||||
"contextLimit": {
|
||||
"type": [
|
||||
"integer",
|
||||
"null"
|
||||
],
|
||||
"format": "uint",
|
||||
"minimum": 0,
|
||||
"description": "Context window size in tokens."
|
||||
},
|
||||
"reasoning": {
|
||||
"type": [
|
||||
"boolean",
|
||||
"null"
|
||||
],
|
||||
"description": "Whether the model supports reasoning/extended thinking."
|
||||
},
|
||||
"recommended": {
|
||||
"type": "boolean",
|
||||
"description": "Whether this model should appear in the compact recommended picker.",
|
||||
"default": false
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"id",
|
||||
"name"
|
||||
],
|
||||
"description": "A single model in provider inventory."
|
||||
},
|
||||
"RefreshProviderInventoryRequest": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"providerIds": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"description": "Which providers to refresh. Empty means all known providers.",
|
||||
"default": []
|
||||
}
|
||||
},
|
||||
"description": "Trigger a background refresh of provider inventories.",
|
||||
"x-side": "agent",
|
||||
"x-method": "_goose/providers/inventory/refresh"
|
||||
},
|
||||
"RefreshProviderInventoryResponse": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"started": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"type": "string"
|
||||
},
|
||||
"description": "Which providers will be refreshed."
|
||||
},
|
||||
"skipped": {
|
||||
"type": "array",
|
||||
"items": {
|
||||
"$ref": "#/$defs/RefreshProviderInventorySkipDto"
|
||||
},
|
||||
"description": "Which providers were skipped and why.",
|
||||
"default": []
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"started"
|
||||
],
|
||||
"description": "Refresh acknowledgement.",
|
||||
"x-side": "agent",
|
||||
"x-method": "_goose/providers/inventory/refresh"
|
||||
},
|
||||
"RefreshProviderInventorySkipDto": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"providerId": {
|
||||
"type": "string"
|
||||
},
|
||||
"reason": {
|
||||
"$ref": "#/$defs/RefreshProviderInventorySkipReasonDto"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"providerId",
|
||||
"reason"
|
||||
]
|
||||
},
|
||||
"RefreshProviderInventorySkipReasonDto": {
|
||||
"type": "string",
|
||||
"enum": [
|
||||
"unknown_provider",
|
||||
"not_configured",
|
||||
"does_not_support_refresh",
|
||||
"already_refreshing"
|
||||
]
|
||||
},
|
||||
"ReadConfigRequest": {
|
||||
"type": "object",
|
||||
@@ -1035,11 +1223,20 @@
|
||||
{
|
||||
"allOf": [
|
||||
{
|
||||
"$ref": "#/$defs/GetProviderModelsRequest"
|
||||
"$ref": "#/$defs/GetProviderInventoryRequest"
|
||||
}
|
||||
],
|
||||
"description": "Params for _goose/providers/models",
|
||||
"title": "GetProviderModelsRequest"
|
||||
"description": "Params for _goose/providers/inventory",
|
||||
"title": "GetProviderInventoryRequest"
|
||||
},
|
||||
{
|
||||
"allOf": [
|
||||
{
|
||||
"$ref": "#/$defs/RefreshProviderInventoryRequest"
|
||||
}
|
||||
],
|
||||
"description": "Params for _goose/providers/inventory/refresh",
|
||||
"title": "RefreshProviderInventoryRequest"
|
||||
},
|
||||
{
|
||||
"allOf": [
|
||||
@@ -1292,10 +1489,18 @@
|
||||
{
|
||||
"allOf": [
|
||||
{
|
||||
"$ref": "#/$defs/GetProviderModelsResponse"
|
||||
"$ref": "#/$defs/GetProviderInventoryResponse"
|
||||
}
|
||||
],
|
||||
"title": "GetProviderModelsResponse"
|
||||
"title": "GetProviderInventoryResponse"
|
||||
},
|
||||
{
|
||||
"allOf": [
|
||||
{
|
||||
"$ref": "#/$defs/RefreshProviderInventoryResponse"
|
||||
}
|
||||
],
|
||||
"title": "RefreshProviderInventoryResponse"
|
||||
},
|
||||
{
|
||||
"allOf": [
|
||||
|
||||
+390
-141
@@ -27,6 +27,9 @@ use goose::mcp_utils::ToolResult;
|
||||
use goose::permission::permission_confirmation::PrincipalType;
|
||||
use goose::permission::{Permission, PermissionConfirmation};
|
||||
use goose::providers::base::Provider;
|
||||
use goose::providers::inventory::{
|
||||
ProviderInventoryEntry, ProviderInventoryService, RefreshSkipReason,
|
||||
};
|
||||
use goose::session::session_manager::SessionType;
|
||||
use goose::session::{EnabledExtensionsState, Session, SessionManager};
|
||||
use goose_acp_macros::custom_methods;
|
||||
@@ -125,6 +128,8 @@ struct AgentSetupRequest {
|
||||
/// Pre-resolved provider name + model config (from config, no network).
|
||||
/// When present the spawn skips re-deriving these from config.
|
||||
resolved_provider: Option<(String, goose::model::ModelConfig)>,
|
||||
/// Pre-instantiated provider reused from synchronous session initialization.
|
||||
prebuilt_provider: Option<Arc<dyn Provider>>,
|
||||
}
|
||||
|
||||
pub struct GooseAcpAgent {
|
||||
@@ -139,6 +144,7 @@ pub struct GooseAcpAgent {
|
||||
permission_manager: Arc<PermissionManager>,
|
||||
goose_mode: GooseMode,
|
||||
disable_session_naming: bool,
|
||||
provider_inventory: ProviderInventoryService,
|
||||
}
|
||||
|
||||
/// Shorten a session/thread id for perf log correlation.
|
||||
@@ -415,19 +421,50 @@ fn builtin_to_extension_config(name: &str) -> ExtensionConfig {
|
||||
}
|
||||
}
|
||||
|
||||
async fn build_model_state(provider: &dyn Provider) -> Result<SessionModelState, sacp::Error> {
|
||||
let models = provider
|
||||
.fetch_recommended_models()
|
||||
.await
|
||||
.map_err(|e| sacp::Error::internal_error().data(e.to_string()))?;
|
||||
let current_model = &provider.get_model_config().model_name;
|
||||
Ok(SessionModelState::new(
|
||||
ModelId::new(current_model.as_str()),
|
||||
models
|
||||
.iter()
|
||||
.map(|name| ModelInfo::new(ModelId::new(&**name), &**name))
|
||||
fn inventory_entry_to_dto(entry: ProviderInventoryEntry) -> ProviderInventoryEntryDto {
|
||||
let stale = ProviderInventoryService::is_stale(&entry);
|
||||
ProviderInventoryEntryDto {
|
||||
provider_id: entry.provider_id,
|
||||
provider_name: entry.provider_name,
|
||||
configured: entry.configured,
|
||||
supports_refresh: entry.supports_refresh,
|
||||
refreshing: entry.refreshing,
|
||||
models: entry
|
||||
.models
|
||||
.into_iter()
|
||||
.map(|m| ProviderInventoryModelDto {
|
||||
id: m.id,
|
||||
name: m.name,
|
||||
family: m.family,
|
||||
context_limit: m.context_limit,
|
||||
reasoning: m.reasoning,
|
||||
recommended: m.recommended,
|
||||
})
|
||||
.collect(),
|
||||
))
|
||||
last_updated_at: entry.last_updated_at.map(|t| t.to_rfc3339()),
|
||||
last_refresh_attempt_at: entry.last_refresh_attempt_at.map(|t| t.to_rfc3339()),
|
||||
last_refresh_error: entry.last_refresh_error,
|
||||
stale,
|
||||
model_selection_hint: entry.model_selection_hint,
|
||||
}
|
||||
}
|
||||
|
||||
fn build_model_state(current_model: &str, inventory: &ProviderInventoryEntry) -> SessionModelState {
|
||||
let mut available_models = inventory
|
||||
.models
|
||||
.iter()
|
||||
.map(|model| ModelInfo::new(ModelId::new(model.id.as_str()), model.name.as_str()))
|
||||
.collect::<Vec<_>>();
|
||||
if !available_models
|
||||
.iter()
|
||||
.any(|model| model.model_id.0.as_ref() == current_model)
|
||||
{
|
||||
available_models.insert(
|
||||
0,
|
||||
ModelInfo::new(ModelId::new(current_model), current_model),
|
||||
);
|
||||
}
|
||||
SessionModelState::new(ModelId::new(current_model), available_models)
|
||||
}
|
||||
|
||||
async fn list_provider_entries(current_provider: Option<&str>) -> Vec<ProviderListEntry> {
|
||||
@@ -546,31 +583,25 @@ fn build_mode_state(current_mode: GooseMode) -> Result<SessionModeState, sacp::E
|
||||
))
|
||||
}
|
||||
|
||||
/// Build model state and config options eagerly from the canonical registry.
|
||||
///
|
||||
/// TODO: This trades speed for correctness — the canonical registry may not perfectly
|
||||
/// match what the provider API returns (new models not yet in the registry, deprecated
|
||||
/// models still listed, or locally-installed models for providers like Ollama). Consider
|
||||
/// whether to reconcile with a live API call in the background.
|
||||
async fn build_eager_config(
|
||||
resolved: &Result<(String, goose::model::ModelConfig), String>,
|
||||
fn should_refresh_inventory_for_session_init(entry: &ProviderInventoryEntry) -> bool {
|
||||
entry.configured
|
||||
&& entry.supports_refresh
|
||||
&& (entry.last_updated_at.is_none() || ProviderInventoryService::is_stale(entry))
|
||||
}
|
||||
|
||||
async fn build_eager_config_from_inventory(
|
||||
provider_name: &str,
|
||||
current_model: &str,
|
||||
inventory: &ProviderInventoryEntry,
|
||||
mode_state: &SessionModeState,
|
||||
goose_session: &Session,
|
||||
) -> (Option<SessionModelState>, Option<Vec<SessionConfigOption>>) {
|
||||
let Ok((ref provider_name, ref mc)) = resolved else {
|
||||
return (None, None);
|
||||
};
|
||||
let recommended = goose::providers::canonical::recommended_models_from_registry(provider_name);
|
||||
let available: Vec<ModelInfo> = recommended
|
||||
.iter()
|
||||
.map(|name| ModelInfo::new(ModelId::new(&**name), &**name))
|
||||
.collect();
|
||||
let ms = SessionModelState::new(ModelId::new(mc.model_name.as_str()), available);
|
||||
) -> (SessionModelState, Vec<SessionConfigOption>) {
|
||||
let ms = build_model_state(current_model, inventory);
|
||||
let provider_selection = session_provider_selection(goose_session);
|
||||
let provider_options = build_provider_options(Some(provider_name.as_str())).await;
|
||||
let provider_options = build_provider_options(Some(provider_name)).await;
|
||||
let config_options =
|
||||
build_config_options(mode_state, &ms, provider_selection, provider_options);
|
||||
(Some(ms), Some(config_options))
|
||||
(ms, config_options)
|
||||
}
|
||||
|
||||
fn build_config_options(
|
||||
@@ -651,6 +682,7 @@ impl GooseAcpAgent {
|
||||
session_manager.storage().clone(),
|
||||
));
|
||||
let permission_manager = Arc::new(PermissionManager::new(config_dir.clone()));
|
||||
let provider_inventory = ProviderInventoryService::new(session_manager.storage().clone());
|
||||
|
||||
Ok(Self {
|
||||
sessions: Arc::new(Mutex::new(HashMap::new())),
|
||||
@@ -664,6 +696,7 @@ impl GooseAcpAgent {
|
||||
permission_manager,
|
||||
goose_mode,
|
||||
disable_session_naming,
|
||||
provider_inventory,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -680,6 +713,125 @@ impl GooseAcpAgent {
|
||||
(self.provider_factory)(provider_name.to_string(), model_config, extensions).await
|
||||
}
|
||||
|
||||
async fn prepare_session_init_config(
|
||||
&self,
|
||||
resolved: &Result<(String, goose::model::ModelConfig), String>,
|
||||
mode_state: &SessionModeState,
|
||||
goose_session: &Session,
|
||||
) -> (
|
||||
Option<SessionModelState>,
|
||||
Option<Vec<SessionConfigOption>>,
|
||||
Option<Arc<dyn Provider>>,
|
||||
) {
|
||||
let Ok((provider_name, model_config)) = resolved else {
|
||||
return (None, None, None);
|
||||
};
|
||||
|
||||
let Some(mut inventory) = self
|
||||
.provider_inventory
|
||||
.entry_for_provider(provider_name)
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
else {
|
||||
return (None, None, None);
|
||||
};
|
||||
|
||||
let mut prebuilt_provider = None;
|
||||
if should_refresh_inventory_for_session_init(&inventory) {
|
||||
match self.load_config() {
|
||||
Ok(config) => {
|
||||
let ext_state = EnabledExtensionsState::extensions_or_default(
|
||||
Some(&goose_session.extension_data),
|
||||
&config,
|
||||
);
|
||||
match self
|
||||
.create_provider(provider_name, model_config.clone(), ext_state)
|
||||
.await
|
||||
{
|
||||
Ok(provider) => {
|
||||
let provider_id = provider_name.clone();
|
||||
prebuilt_provider = Some(provider.clone());
|
||||
match self
|
||||
.provider_inventory
|
||||
.plan_refresh(std::slice::from_ref(&provider_id))
|
||||
.await
|
||||
{
|
||||
Ok(plan) if plan.started.iter().any(|id| id == &provider_id) => {
|
||||
match provider.fetch_recommended_models().await {
|
||||
Ok(models) => {
|
||||
if let Err(error) = self
|
||||
.provider_inventory
|
||||
.store_refreshed_models(&provider_id, &models)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
provider = %provider_id,
|
||||
error = %error,
|
||||
"failed to store refreshed provider inventory during session init"
|
||||
);
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
if let Err(store_error) = self
|
||||
.provider_inventory
|
||||
.store_refresh_error(
|
||||
&provider_id,
|
||||
error.to_string(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
provider = %provider_id,
|
||||
error = %store_error,
|
||||
"failed to store provider inventory refresh error during session init"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(_) => {}
|
||||
Err(error) => warn!(
|
||||
provider = %provider_id,
|
||||
error = %error,
|
||||
"failed to plan provider inventory refresh during session init"
|
||||
),
|
||||
}
|
||||
|
||||
if let Ok(Some(refreshed_inventory)) = self
|
||||
.provider_inventory
|
||||
.entry_for_provider(provider_name)
|
||||
.await
|
||||
{
|
||||
inventory = refreshed_inventory;
|
||||
}
|
||||
}
|
||||
Err(error) => warn!(
|
||||
provider = %provider_name,
|
||||
error = %error,
|
||||
"failed to initialize provider during synchronous inventory refresh"
|
||||
),
|
||||
}
|
||||
}
|
||||
Err(error) => warn!(
|
||||
provider = %provider_name,
|
||||
error = %error,
|
||||
"failed to load config during synchronous inventory refresh"
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
let (model_state, config_options) = build_eager_config_from_inventory(
|
||||
provider_name,
|
||||
model_config.model_name.as_str(),
|
||||
&inventory,
|
||||
mode_state,
|
||||
goose_session,
|
||||
)
|
||||
.await;
|
||||
(Some(model_state), Some(config_options), prebuilt_provider)
|
||||
}
|
||||
|
||||
fn spawn_agent_setup(
|
||||
&self,
|
||||
cx: &ConnectionTo<Client>,
|
||||
@@ -691,6 +843,7 @@ impl GooseAcpAgent {
|
||||
goose_session,
|
||||
mcp_servers,
|
||||
resolved_provider,
|
||||
prebuilt_provider,
|
||||
} = req;
|
||||
|
||||
let goose_mode = goose_session.goose_mode;
|
||||
@@ -845,9 +998,12 @@ impl GooseAcpAgent {
|
||||
Some(&goose_session.extension_data),
|
||||
&config,
|
||||
);
|
||||
let provider = provider_factory(provider_name.to_string(), model_config, ext_state)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?;
|
||||
let provider = match prebuilt_provider {
|
||||
Some(provider) => provider,
|
||||
None => provider_factory(provider_name.to_string(), model_config, ext_state)
|
||||
.await
|
||||
.map_err(|e| e.to_string())?,
|
||||
};
|
||||
agent
|
||||
.update_provider(provider.clone(), &goose_session.id)
|
||||
.await
|
||||
@@ -1416,9 +1572,10 @@ impl GooseAcpAgent {
|
||||
.as_ref()
|
||||
.ok()
|
||||
.map(|(_, mc)| build_usage_update(&goose_session, mc.context_limit()));
|
||||
let (model_state, config_options) =
|
||||
build_eager_config(&resolved, &mode_state, &goose_session).await;
|
||||
let session_id = SessionId::new(thread_id.clone());
|
||||
let (model_state, config_options, prebuilt_provider) = self
|
||||
.prepare_session_init_config(&resolved, &mode_state, &goose_session)
|
||||
.await;
|
||||
|
||||
self.spawn_agent_setup(
|
||||
cx,
|
||||
@@ -1428,6 +1585,7 @@ impl GooseAcpAgent {
|
||||
goose_session,
|
||||
mcp_servers: args.mcp_servers,
|
||||
resolved_provider: resolved.as_ref().ok().cloned(),
|
||||
prebuilt_provider,
|
||||
},
|
||||
);
|
||||
|
||||
@@ -1798,8 +1956,9 @@ impl GooseAcpAgent {
|
||||
.as_ref()
|
||||
.map(|mc| build_usage_update(&goose_session, mc.context_limit()))
|
||||
});
|
||||
let (model_state, config_options) =
|
||||
build_eager_config(&resolved, &mode_state, &goose_session).await;
|
||||
let (model_state, config_options, prebuilt_provider) = self
|
||||
.prepare_session_init_config(&resolved, &mode_state, &goose_session)
|
||||
.await;
|
||||
|
||||
self.spawn_agent_setup(
|
||||
cx,
|
||||
@@ -1809,6 +1968,7 @@ impl GooseAcpAgent {
|
||||
goose_session,
|
||||
mcp_servers: args.mcp_servers,
|
||||
resolved_provider: None,
|
||||
prebuilt_provider,
|
||||
},
|
||||
);
|
||||
|
||||
@@ -2116,10 +2276,21 @@ impl GooseAcpAgent {
|
||||
let provider = agent.provider().await.map_err(|e| {
|
||||
sacp::Error::internal_error().data(format!("Failed to get provider: {}", e))
|
||||
})?;
|
||||
let provider_name = provider.get_name().to_string();
|
||||
let current_model = provider.get_model_config().model_name.clone();
|
||||
let goose_mode = agent.goose_mode().await;
|
||||
let model_state = build_model_state(&*provider).await?;
|
||||
let inventory = self
|
||||
.provider_inventory
|
||||
.entry_for_provider(&provider_name)
|
||||
.await
|
||||
.map_err(|e| sacp::Error::internal_error().data(e.to_string()))?;
|
||||
let Some(inventory) = inventory else {
|
||||
return Err(sacp::Error::internal_error()
|
||||
.data(format!("Unknown provider inventory: {}", provider_name)));
|
||||
};
|
||||
let model_state = build_model_state(current_model.as_str(), &inventory);
|
||||
let mode_state = build_mode_state(goose_mode)?;
|
||||
let provider_options = build_provider_options(Some(provider.get_name())).await;
|
||||
let provider_options = build_provider_options(Some(&provider_name)).await;
|
||||
let config_options = build_config_options(
|
||||
&mode_state,
|
||||
&model_state,
|
||||
@@ -2399,8 +2570,9 @@ impl GooseAcpAgent {
|
||||
|
||||
let mode_state = build_mode_state(self.goose_mode)?;
|
||||
let resolved = resolve_provider_and_model(&self.config_dir, &goose_session).await;
|
||||
let (model_state, config_options) =
|
||||
build_eager_config(&resolved, &mode_state, &goose_session).await;
|
||||
let (model_state, config_options, prebuilt_provider) = self
|
||||
.prepare_session_init_config(&resolved, &mode_state, &goose_session)
|
||||
.await;
|
||||
|
||||
self.spawn_agent_setup(
|
||||
cx,
|
||||
@@ -2410,6 +2582,7 @@ impl GooseAcpAgent {
|
||||
goose_session,
|
||||
mcp_servers: args.mcp_servers,
|
||||
resolved_provider: resolved.ok(),
|
||||
prebuilt_provider,
|
||||
},
|
||||
);
|
||||
|
||||
@@ -2677,58 +2850,82 @@ impl GooseAcpAgent {
|
||||
Ok(GetProviderDetailsResponse { providers: entries })
|
||||
}
|
||||
|
||||
#[custom_method(GetProviderModelsRequest)]
|
||||
async fn on_get_provider_models(
|
||||
#[custom_method(GetProviderInventoryRequest)]
|
||||
async fn on_get_provider_inventory(
|
||||
&self,
|
||||
req: GetProviderModelsRequest,
|
||||
) -> Result<GetProviderModelsResponse, sacp::Error> {
|
||||
let config = self.load_config().ok();
|
||||
let all = goose::providers::providers().await;
|
||||
req: GetProviderInventoryRequest,
|
||||
) -> Result<GetProviderInventoryResponse, sacp::Error> {
|
||||
let entries = self
|
||||
.provider_inventory
|
||||
.entries(&req.provider_ids)
|
||||
.await
|
||||
.map_err(|e| sacp::Error::internal_error().data(e.to_string()))?;
|
||||
Ok(GetProviderInventoryResponse {
|
||||
entries: entries.into_iter().map(inventory_entry_to_dto).collect(),
|
||||
})
|
||||
}
|
||||
|
||||
let Some((metadata, _provider_type)) =
|
||||
all.into_iter().find(|(m, _)| m.name == req.provider_name)
|
||||
else {
|
||||
return Err(sacp::Error::invalid_params()
|
||||
.data(format!("Unknown provider: {}", req.provider_name)));
|
||||
};
|
||||
|
||||
let is_configured = config
|
||||
.as_ref()
|
||||
.map(|c| {
|
||||
metadata.config_keys.iter().all(|k| {
|
||||
if !k.required {
|
||||
return true;
|
||||
}
|
||||
if k.secret {
|
||||
c.get_secret::<String>(&k.name).is_ok()
|
||||
} else {
|
||||
c.get_param::<String>(&k.name).is_ok()
|
||||
}
|
||||
})
|
||||
})
|
||||
.unwrap_or(false);
|
||||
|
||||
if !is_configured {
|
||||
return Err(sacp::Error::invalid_params().data(format!(
|
||||
"Provider '{}' is not configured",
|
||||
req.provider_name
|
||||
)));
|
||||
#[custom_method(RefreshProviderInventoryRequest)]
|
||||
async fn on_refresh_provider_inventory(
|
||||
&self,
|
||||
req: RefreshProviderInventoryRequest,
|
||||
) -> Result<RefreshProviderInventoryResponse, sacp::Error> {
|
||||
let refresh_plan = self
|
||||
.provider_inventory
|
||||
.plan_refresh(&req.provider_ids)
|
||||
.await;
|
||||
let refresh_plan =
|
||||
refresh_plan.map_err(|e| sacp::Error::internal_error().data(e.to_string()))?;
|
||||
for provider_id in &refresh_plan.started {
|
||||
let provider_inventory = self.provider_inventory.clone();
|
||||
let provider_factory = Arc::clone(&self.provider_factory);
|
||||
let provider_id = provider_id.clone();
|
||||
tokio::spawn(async move {
|
||||
let result = async {
|
||||
let metadata = goose::providers::get_from_registry(&provider_id).await?;
|
||||
let model_config =
|
||||
goose::model::ModelConfig::new(&metadata.metadata().default_model)?
|
||||
.with_canonical_limits(&provider_id);
|
||||
let provider =
|
||||
provider_factory(provider_id.clone(), model_config, Vec::new()).await?;
|
||||
let models = provider.fetch_recommended_models().await?;
|
||||
provider_inventory
|
||||
.store_refreshed_models(&provider_id, &models)
|
||||
.await
|
||||
}
|
||||
.await;
|
||||
if let Err(error) = result {
|
||||
let _ = provider_inventory
|
||||
.store_refresh_error(&provider_id, error.to_string())
|
||||
.await;
|
||||
warn!(provider = %provider_id, error = %error, "provider inventory refresh failed");
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
let model_config = goose::model::ModelConfig::new(&metadata.default_model)
|
||||
.map_err(|e| sacp::Error::internal_error().data(e.to_string()))?
|
||||
.with_canonical_limits(&req.provider_name);
|
||||
|
||||
let provider = (self.provider_factory)(req.provider_name.clone(), model_config, Vec::new())
|
||||
.await
|
||||
.map_err(|e| sacp::Error::internal_error().data(e.to_string()))?;
|
||||
|
||||
let models = provider
|
||||
.fetch_recommended_models()
|
||||
.await
|
||||
.map_err(|e| sacp::Error::internal_error().data(e.to_string()))?;
|
||||
|
||||
Ok(GetProviderModelsResponse { models })
|
||||
Ok(RefreshProviderInventoryResponse {
|
||||
started: refresh_plan.started,
|
||||
skipped: refresh_plan
|
||||
.skipped
|
||||
.into_iter()
|
||||
.map(|entry| RefreshProviderInventorySkipDto {
|
||||
provider_id: entry.provider_id,
|
||||
reason: match entry.reason {
|
||||
RefreshSkipReason::UnknownProvider => {
|
||||
RefreshProviderInventorySkipReasonDto::UnknownProvider
|
||||
}
|
||||
RefreshSkipReason::NotConfigured => {
|
||||
RefreshProviderInventorySkipReasonDto::NotConfigured
|
||||
}
|
||||
RefreshSkipReason::DoesNotSupportRefresh => {
|
||||
RefreshProviderInventorySkipReasonDto::DoesNotSupportRefresh
|
||||
}
|
||||
RefreshSkipReason::AlreadyRefreshing => {
|
||||
RefreshProviderInventorySkipReasonDto::AlreadyRefreshing
|
||||
}
|
||||
},
|
||||
})
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
|
||||
#[custom_method(ReadConfigRequest)]
|
||||
@@ -3450,11 +3647,77 @@ impl HandleDispatchFrom<Client> for GooseAcpHandler {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
// Respond immediately using the current provider inventory snapshot.
|
||||
let t_tail = std::time::Instant::now();
|
||||
let (notification, config_options) = agent.build_config_update(&session_id).await?;
|
||||
cx.send_notification(notification)?;
|
||||
responder.respond(SetSessionConfigOptionResponse::new(config_options))?;
|
||||
debug!(target: "perf", sid = %sid, ms = t_tail.elapsed().as_millis() as u64, "perf: set_config_option notification_and_respond");
|
||||
debug!(target: "perf", sid = %sid, ms = t_tail.elapsed().as_millis() as u64, "perf: set_config_option inventory_respond");
|
||||
|
||||
let maybe_refresh = if config_id == "provider" {
|
||||
let provider_id = value_id.0.to_string();
|
||||
agent
|
||||
.provider_inventory
|
||||
.plan_refresh(std::slice::from_ref(&provider_id))
|
||||
.await
|
||||
.ok()
|
||||
.filter(|plan| plan.started.iter().any(|id| id == &provider_id))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if maybe_refresh.is_some() {
|
||||
let agent_bg = agent.clone();
|
||||
let cx_bg = cx.clone();
|
||||
let session_id_bg = session_id.clone();
|
||||
let sid_bg = sid.clone();
|
||||
tokio::spawn(async move {
|
||||
let t_bg = std::time::Instant::now();
|
||||
let refreshed = async {
|
||||
let session_agent =
|
||||
agent_bg.get_session_agent(&session_id_bg.0, None).await?;
|
||||
let provider = session_agent
|
||||
.provider()
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!(e.to_string()))?;
|
||||
let provider_name = provider.get_name().to_string();
|
||||
let models = provider
|
||||
.fetch_recommended_models()
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!(e.to_string()))?;
|
||||
agent_bg
|
||||
.provider_inventory
|
||||
.store_refreshed_models(&provider_name, &models)
|
||||
.await?;
|
||||
agent_bg
|
||||
.build_config_update(&session_id_bg)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!(e.to_string()))
|
||||
}
|
||||
.await;
|
||||
|
||||
match refreshed {
|
||||
Ok((fresh_notification, _)) => {
|
||||
let _ = cx_bg.send_notification(fresh_notification);
|
||||
debug!(target: "perf", sid = %sid_bg, ms = t_bg.elapsed().as_millis() as u64, "perf: set_config_option background_refresh done");
|
||||
}
|
||||
Err(e) => {
|
||||
if let Ok(session_agent) =
|
||||
agent_bg.get_session_agent(&session_id_bg.0, None).await
|
||||
{
|
||||
if let Ok(provider) = session_agent.provider().await {
|
||||
let provider_name = provider.get_name().to_string();
|
||||
let _ = agent_bg
|
||||
.provider_inventory
|
||||
.store_refresh_error(&provider_name, e.to_string())
|
||||
.await;
|
||||
}
|
||||
}
|
||||
debug!(target: "perf", sid = %sid_bg, error = %e, ms = t_bg.elapsed().as_millis() as u64, "perf: set_config_option background_refresh failed");
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
debug!(target: "perf", sid = %sid, ms = t_handler.elapsed().as_millis() as u64, config_id = %config_id, "perf: set_config_option done");
|
||||
Ok(())
|
||||
}
|
||||
@@ -3597,7 +3860,6 @@ pub async fn run(builtins: Vec<String>) -> Result<()> {
|
||||
mod tests {
|
||||
use super::*;
|
||||
use goose::conversation::message::{ToolRequest, ToolResponse};
|
||||
use goose::providers::errors::ProviderError;
|
||||
use rmcp::model::{CallToolRequestParams, Content as RmcpContent};
|
||||
use sacp::schema::{
|
||||
EnvVariable, HttpHeader, McpServer, McpServerHttp, McpServerSse, McpServerStdio,
|
||||
@@ -3787,61 +4049,48 @@ print(\"hello, world\")
|
||||
assert_eq!(outcome_to_confirmation(&input), expected);
|
||||
}
|
||||
|
||||
struct MockModelProvider {
|
||||
models: Result<Vec<String>, ProviderError>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl Provider for MockModelProvider {
|
||||
fn get_name(&self) -> &str {
|
||||
"mock"
|
||||
}
|
||||
|
||||
async fn stream(
|
||||
&self,
|
||||
_model_config: &goose::model::ModelConfig,
|
||||
_session_id: &str,
|
||||
_system: &str,
|
||||
_messages: &[goose::conversation::message::Message],
|
||||
_tools: &[rmcp::model::Tool],
|
||||
) -> Result<goose::providers::base::MessageStream, ProviderError> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
fn get_model_config(&self) -> goose::model::ModelConfig {
|
||||
goose::model::ModelConfig::new_or_fail("unused")
|
||||
}
|
||||
|
||||
async fn fetch_recommended_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||
self.models.clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[test_case(
|
||||
Ok(vec!["model-a".into(), "model-b".into()])
|
||||
=> Ok(SessionModelState::new(
|
||||
vec!["model-a".into(), "model-b".into()]
|
||||
=> SessionModelState::new(
|
||||
ModelId::new("unused"),
|
||||
vec![ModelInfo::new(ModelId::new("model-a"), "model-a"),
|
||||
vec![ModelInfo::new(ModelId::new("unused"), "unused"),
|
||||
ModelInfo::new(ModelId::new("model-a"), "model-a"),
|
||||
ModelInfo::new(ModelId::new("model-b"), "model-b")],
|
||||
))
|
||||
)
|
||||
; "returns current and available models"
|
||||
)]
|
||||
#[test_case(
|
||||
Ok(vec![])
|
||||
=> Ok(SessionModelState::new(ModelId::new("unused"), vec![]))
|
||||
vec![]
|
||||
=> SessionModelState::new(
|
||||
ModelId::new("unused"),
|
||||
vec![ModelInfo::new(ModelId::new("unused"), "unused")],
|
||||
)
|
||||
; "empty model list"
|
||||
)]
|
||||
#[test_case(
|
||||
Err(ProviderError::ExecutionError("fail".into()))
|
||||
=> Err(sacp::Error::internal_error().data("Execution error: fail".to_string()))
|
||||
; "fetch error propagates"
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn test_build_model_state(
|
||||
models: Result<Vec<String>, ProviderError>,
|
||||
) -> Result<SessionModelState, sacp::Error> {
|
||||
let provider = MockModelProvider { models };
|
||||
build_model_state(&provider).await
|
||||
fn test_build_model_state(models: Vec<String>) -> SessionModelState {
|
||||
let inventory = ProviderInventoryEntry {
|
||||
provider_id: "mock".to_string(),
|
||||
provider_name: "Mock".to_string(),
|
||||
configured: true,
|
||||
supports_refresh: true,
|
||||
refreshing: false,
|
||||
models: models
|
||||
.into_iter()
|
||||
.map(|id| goose::providers::inventory::InventoryModel {
|
||||
name: id.clone(),
|
||||
id,
|
||||
family: None,
|
||||
context_limit: None,
|
||||
reasoning: None,
|
||||
recommended: false,
|
||||
})
|
||||
.collect(),
|
||||
last_updated_at: None,
|
||||
last_refresh_attempt_at: None,
|
||||
last_refresh_error: None,
|
||||
model_selection_hint: None,
|
||||
};
|
||||
build_model_state("unused", &inventory)
|
||||
}
|
||||
|
||||
fn json_object(pairs: Vec<(&str, serde_json::Value)>) -> rmcp::model::JsonObject {
|
||||
|
||||
@@ -257,20 +257,6 @@ pub struct GetProviderDetailsResponse {
|
||||
pub providers: Vec<ProviderDetailEntry>,
|
||||
}
|
||||
|
||||
/// Fetch the full list of models available for a specific provider.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
|
||||
#[request(method = "_goose/providers/models", response = GetProviderModelsResponse)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct GetProviderModelsRequest {
|
||||
pub provider_name: String,
|
||||
}
|
||||
|
||||
/// Provider models response.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
|
||||
pub struct GetProviderModelsResponse {
|
||||
pub models: Vec<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ProviderDetailEntry {
|
||||
@@ -370,6 +356,120 @@ pub struct DictationConfigResponse {
|
||||
pub providers: HashMap<String, DictationProviderStatusEntry>,
|
||||
}
|
||||
|
||||
/// Read per-provider inventory. Always returns immediately from stored state.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
|
||||
#[request(
|
||||
method = "_goose/providers/inventory",
|
||||
response = GetProviderInventoryResponse
|
||||
)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct GetProviderInventoryRequest {
|
||||
/// Only return entries for these providers. Empty means all.
|
||||
#[serde(default)]
|
||||
pub provider_ids: Vec<String>,
|
||||
}
|
||||
|
||||
/// Provider inventory response.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
|
||||
pub struct GetProviderInventoryResponse {
|
||||
pub entries: Vec<ProviderInventoryEntryDto>,
|
||||
}
|
||||
|
||||
/// Trigger a background refresh of provider inventories.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
|
||||
#[request(
|
||||
method = "_goose/providers/inventory/refresh",
|
||||
response = RefreshProviderInventoryResponse
|
||||
)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct RefreshProviderInventoryRequest {
|
||||
/// Which providers to refresh. Empty means all known providers.
|
||||
#[serde(default)]
|
||||
pub provider_ids: Vec<String>,
|
||||
}
|
||||
|
||||
/// Refresh acknowledgement.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct RefreshProviderInventoryResponse {
|
||||
/// Which providers will be refreshed.
|
||||
pub started: Vec<String>,
|
||||
/// Which providers were skipped and why.
|
||||
#[serde(default)]
|
||||
pub skipped: Vec<RefreshProviderInventorySkipDto>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct RefreshProviderInventorySkipDto {
|
||||
pub provider_id: String,
|
||||
pub reason: RefreshProviderInventorySkipReasonDto,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum RefreshProviderInventorySkipReasonDto {
|
||||
#[default]
|
||||
UnknownProvider,
|
||||
NotConfigured,
|
||||
DoesNotSupportRefresh,
|
||||
AlreadyRefreshing,
|
||||
}
|
||||
|
||||
/// A single model in provider inventory.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ProviderInventoryModelDto {
|
||||
/// Model identifier as the provider knows it.
|
||||
pub id: String,
|
||||
/// Human-readable display name.
|
||||
pub name: String,
|
||||
/// Model family for grouping in UI.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub family: Option<String>,
|
||||
/// Context window size in tokens.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub context_limit: Option<usize>,
|
||||
/// Whether the model supports reasoning/extended thinking.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning: Option<bool>,
|
||||
/// Whether this model should appear in the compact recommended picker.
|
||||
#[serde(default)]
|
||||
pub recommended: bool,
|
||||
}
|
||||
|
||||
/// Provider inventory entry.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ProviderInventoryEntryDto {
|
||||
/// Provider identifier.
|
||||
pub provider_id: String,
|
||||
/// Human-readable provider name.
|
||||
pub provider_name: String,
|
||||
/// Whether Goose has enough configuration to use this provider.
|
||||
pub configured: bool,
|
||||
/// Whether this provider supports background inventory refresh.
|
||||
pub supports_refresh: bool,
|
||||
/// Whether a refresh is currently in flight.
|
||||
pub refreshing: bool,
|
||||
/// The list of available models.
|
||||
pub models: Vec<ProviderInventoryModelDto>,
|
||||
/// When this entry was last successfully refreshed (ISO 8601).
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub last_updated_at: Option<String>,
|
||||
/// When a refresh was most recently attempted (ISO 8601).
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub last_refresh_attempt_at: Option<String>,
|
||||
/// The last refresh failure message, if any.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub last_refresh_error: Option<String>,
|
||||
/// Whether we believe this data may be outdated.
|
||||
pub stale: bool,
|
||||
/// Guidance message shown when this provider manages its own model selection externally.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub model_selection_hint: Option<String>,
|
||||
}
|
||||
|
||||
/// Empty success response for operations that return no data.
|
||||
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
|
||||
pub struct EmptyResponse {}
|
||||
|
||||
@@ -416,6 +416,11 @@ mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
let tmp_dir = tempfile::tempdir().unwrap();
|
||||
let temp_root = tmp_dir.path().display().to_string();
|
||||
let _guard = env_lock::lock_env([
|
||||
("HOME", Some(temp_root.as_str())),
|
||||
("GOOSE_PATH_ROOT", Some(temp_root.as_str())),
|
||||
]);
|
||||
let session_manager = Arc::new(SessionManager::new(tmp_dir.path().to_path_buf()));
|
||||
let session = session_manager
|
||||
.create_session(
|
||||
|
||||
@@ -2,6 +2,7 @@ use crate::config::paths::Paths;
|
||||
use crate::config::Config;
|
||||
use crate::providers::anthropic::AnthropicProvider;
|
||||
use crate::providers::base::{ModelInfo, ProviderType};
|
||||
use crate::providers::inventory::declarative_inventory_identity;
|
||||
use crate::providers::ollama::OllamaProvider;
|
||||
use crate::providers::openai::OpenAiProvider;
|
||||
use anyhow::Result;
|
||||
@@ -460,38 +461,59 @@ pub fn register_declarative_provider(
|
||||
match config.engine {
|
||||
ProviderEngine::OpenAI => {
|
||||
let captured = config.clone();
|
||||
registry.register_with_name::<OpenAiProvider, _>(
|
||||
let identity_config = config.clone();
|
||||
registry.register_with_name::<OpenAiProvider, _, _>(
|
||||
&config,
|
||||
provider_type,
|
||||
config.dynamic_models.unwrap_or(false),
|
||||
move |model| {
|
||||
let mut cfg = captured.clone();
|
||||
resolve_config(&mut cfg)?;
|
||||
OpenAiProvider::from_custom_config(model, cfg)
|
||||
},
|
||||
move || {
|
||||
let mut cfg = identity_config.clone();
|
||||
resolve_config(&mut cfg)?;
|
||||
declarative_inventory_identity(&cfg)
|
||||
},
|
||||
);
|
||||
}
|
||||
ProviderEngine::Ollama => {
|
||||
let captured = config.clone();
|
||||
registry.register_with_name::<OllamaProvider, _>(
|
||||
let identity_config = config.clone();
|
||||
registry.register_with_name::<OllamaProvider, _, _>(
|
||||
&config,
|
||||
provider_type,
|
||||
config.dynamic_models.unwrap_or(false),
|
||||
move |model| {
|
||||
let mut cfg = captured.clone();
|
||||
resolve_config(&mut cfg)?;
|
||||
OllamaProvider::from_custom_config(model, cfg)
|
||||
},
|
||||
move || {
|
||||
let mut cfg = identity_config.clone();
|
||||
resolve_config(&mut cfg)?;
|
||||
declarative_inventory_identity(&cfg)
|
||||
},
|
||||
);
|
||||
}
|
||||
ProviderEngine::Anthropic => {
|
||||
let captured = config.clone();
|
||||
registry.register_with_name::<AnthropicProvider, _>(
|
||||
let identity_config = config.clone();
|
||||
registry.register_with_name::<AnthropicProvider, _, _>(
|
||||
&config,
|
||||
provider_type,
|
||||
config.dynamic_models.unwrap_or(false),
|
||||
move |model| {
|
||||
let mut cfg = captured.clone();
|
||||
resolve_config(&mut cfg)?;
|
||||
AnthropicProvider::from_custom_config(model, cfg)
|
||||
},
|
||||
move || {
|
||||
let mut cfg = identity_config.clone();
|
||||
resolve_config(&mut cfg)?;
|
||||
declarative_inventory_identity(&cfg)
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
use crate::config::search_path::SearchPaths;
|
||||
use crate::providers::inventory::InventoryIdentityInput;
|
||||
use anyhow::Result;
|
||||
use std::path::PathBuf;
|
||||
|
||||
pub fn acp_adapter_installed(command: &str) -> bool {
|
||||
resolve_acp_command(command).is_ok()
|
||||
}
|
||||
|
||||
pub fn acp_inventory_identity(provider_id: &str, command: &str) -> Result<InventoryIdentityInput> {
|
||||
let resolved_command = resolve_acp_command(command)?;
|
||||
Ok(InventoryIdentityInput::new(provider_id, provider_id)
|
||||
.with_public("command", resolved_command.display().to_string()))
|
||||
}
|
||||
|
||||
fn resolve_acp_command(command: &str) -> Result<PathBuf> {
|
||||
SearchPaths::builder().with_npm().resolve(command)
|
||||
}
|
||||
@@ -9,7 +9,9 @@ use crate::acp::{
|
||||
use crate::config::search_path::SearchPaths;
|
||||
use crate::config::{Config, GooseMode};
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity};
|
||||
use crate::providers::base::{ProviderDef, ProviderMetadata};
|
||||
use crate::providers::inventory::InventoryIdentityInput;
|
||||
|
||||
const AMP_ACP_PROVIDER_NAME: &str = "amp-acp";
|
||||
const AMP_ACP_DOC_URL: &str = "https://ampcode.com";
|
||||
@@ -37,6 +39,7 @@ impl ProviderDef for AmpAcpProvider {
|
||||
"Set in your goose config file (`~/.config/goose/config.yaml` on macOS/Linux):\n GOOSE_PROVIDER: amp-acp\n GOOSE_MODEL: current",
|
||||
"Restart goose for changes to take effect",
|
||||
])
|
||||
.with_model_selection_hint("Use the Amp CLI to configure models")
|
||||
}
|
||||
|
||||
fn from_env(
|
||||
@@ -49,10 +52,12 @@ impl ProviderDef for AmpAcpProvider {
|
||||
let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto);
|
||||
|
||||
let mode_mapping = HashMap::from([
|
||||
(GooseMode::Auto, "auto".to_string()),
|
||||
(GooseMode::Approve, "approve".to_string()),
|
||||
(GooseMode::SmartApprove, "smart-approve".to_string()),
|
||||
(GooseMode::Chat, "chat".to_string()),
|
||||
// "bypass" skips confirmations, closest to autonomous mode.
|
||||
(GooseMode::Auto, "bypass".to_string()),
|
||||
// "default" prompts before risky actions.
|
||||
(GooseMode::Approve, "default".to_string()),
|
||||
(GooseMode::SmartApprove, "default".to_string()),
|
||||
(GooseMode::Chat, "default".to_string()),
|
||||
]);
|
||||
|
||||
let provider_config = AcpProviderConfig {
|
||||
@@ -71,4 +76,16 @@ impl ProviderDef for AmpAcpProvider {
|
||||
AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await
|
||||
})
|
||||
}
|
||||
|
||||
fn supports_inventory_refresh() -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn inventory_identity() -> Result<InventoryIdentityInput> {
|
||||
acp_inventory_identity(AMP_ACP_PROVIDER_NAME, AMP_ACP_BINARY)
|
||||
}
|
||||
|
||||
fn inventory_configured() -> bool {
|
||||
acp_adapter_installed(AMP_ACP_BINARY)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,6 +14,7 @@ use super::errors::ProviderError;
|
||||
use super::formats::anthropic::{
|
||||
create_request, response_to_streaming_message, thinking_type, ThinkingType,
|
||||
};
|
||||
use super::inventory::{config_secret_value, serialize_string_map, InventoryIdentityInput};
|
||||
use super::openai_compatible::handle_status_openai_compat;
|
||||
use super::openai_compatible::map_http_error_to_provider_error;
|
||||
use super::retry::ProviderRetry;
|
||||
@@ -235,6 +236,33 @@ impl ProviderDef for AnthropicProvider {
|
||||
) -> BoxFuture<'static, Result<Self::Provider>> {
|
||||
Box::pin(Self::from_env(model))
|
||||
}
|
||||
|
||||
fn supports_inventory_refresh() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn inventory_identity() -> Result<InventoryIdentityInput> {
|
||||
let config = crate::config::Config::global();
|
||||
let mut identity =
|
||||
InventoryIdentityInput::new(ANTHROPIC_PROVIDER_NAME, ANTHROPIC_PROVIDER_NAME)
|
||||
.with_public(
|
||||
"host",
|
||||
config
|
||||
.get_param::<String>("ANTHROPIC_HOST")
|
||||
.unwrap_or_else(|_| "https://api.anthropic.com".to_string()),
|
||||
);
|
||||
|
||||
if let Some(api_key) = config_secret_value(config, "ANTHROPIC_API_KEY") {
|
||||
identity = identity.with_secret("api_key", api_key);
|
||||
}
|
||||
if let Ok(headers) = config
|
||||
.get_secret::<std::collections::HashMap<String, String>>("ANTHROPIC_CUSTOM_HEADERS")
|
||||
{
|
||||
identity = identity.with_secret("headers", serialize_string_map(&headers)?);
|
||||
}
|
||||
|
||||
Ok(identity)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
|
||||
@@ -6,9 +6,10 @@ use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::canonical::{map_to_canonical_model, CanonicalModelRegistry};
|
||||
use super::errors::ProviderError;
|
||||
use super::inventory::{default_inventory_identity, InventoryIdentityInput};
|
||||
use super::retry::RetryConfig;
|
||||
use crate::config::base::ConfigValue;
|
||||
use crate::config::{ExtensionConfig, GooseMode};
|
||||
use crate::config::{Config, ExtensionConfig, GooseMode};
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::conversation::Conversation;
|
||||
use crate::model::ModelConfig;
|
||||
@@ -179,6 +180,9 @@ pub struct ProviderMetadata {
|
||||
/// step-by-step instructions for set up providers eg: api key
|
||||
#[serde(default)]
|
||||
pub setup_steps: Vec<String>,
|
||||
/// Hint shown in the model picker when this provider manages its own model selection.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub model_selection_hint: Option<String>,
|
||||
}
|
||||
|
||||
impl ProviderMetadata {
|
||||
@@ -212,6 +216,7 @@ impl ProviderMetadata {
|
||||
model_doc_link: model_doc_link.to_string(),
|
||||
config_keys,
|
||||
setup_steps: vec![],
|
||||
model_selection_hint: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -233,6 +238,7 @@ impl ProviderMetadata {
|
||||
model_doc_link: model_doc_link.to_string(),
|
||||
config_keys,
|
||||
setup_steps: vec![],
|
||||
model_selection_hint: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -246,6 +252,7 @@ impl ProviderMetadata {
|
||||
model_doc_link: "".to_string(),
|
||||
config_keys: vec![],
|
||||
setup_steps: vec![],
|
||||
model_selection_hint: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -253,6 +260,11 @@ impl ProviderMetadata {
|
||||
self.setup_steps = steps.into_iter().map(|s| s.to_string()).collect();
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_model_selection_hint(mut self, hint: &str) -> Self {
|
||||
self.model_selection_hint = Some(hint.to_string());
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// Configuration key metadata for provider setup
|
||||
@@ -492,6 +504,34 @@ pub trait ProviderDef: Send + Sync {
|
||||
) -> BoxFuture<'static, Result<Self::Provider>>
|
||||
where
|
||||
Self: Sized;
|
||||
|
||||
fn supports_inventory_refresh() -> bool
|
||||
where
|
||||
Self: Sized,
|
||||
{
|
||||
false
|
||||
}
|
||||
|
||||
fn inventory_identity() -> Result<InventoryIdentityInput>
|
||||
where
|
||||
Self: Sized,
|
||||
{
|
||||
let metadata = Self::metadata();
|
||||
Ok(default_inventory_identity(
|
||||
&metadata.name,
|
||||
&metadata.name,
|
||||
&metadata.config_keys,
|
||||
Config::global(),
|
||||
))
|
||||
}
|
||||
|
||||
fn inventory_configured() -> bool
|
||||
where
|
||||
Self: Sized,
|
||||
{
|
||||
let metadata = Self::metadata();
|
||||
super::inventory::default_inventory_configured(&metadata.config_keys, Config::global())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
@@ -588,7 +628,7 @@ pub trait Provider: Send + Sync {
|
||||
false
|
||||
}
|
||||
|
||||
/// Fetch models filtered by canonical registry and usability
|
||||
/// Fetch inventory models filtered by canonical registry and usability.
|
||||
async fn fetch_recommended_models(&self) -> Result<Vec<String>, ProviderError> {
|
||||
let all_models = self.fetch_supported_models().await?;
|
||||
|
||||
@@ -637,15 +677,15 @@ pub trait Provider: Send + Sync {
|
||||
(None, None) => a.0.cmp(&b.0),
|
||||
});
|
||||
|
||||
let recommended_models: Vec<String> = models_with_dates
|
||||
let inventory_models: Vec<String> = models_with_dates
|
||||
.into_iter()
|
||||
.map(|(name, _)| name)
|
||||
.collect();
|
||||
|
||||
if recommended_models.is_empty() {
|
||||
if inventory_models.is_empty() {
|
||||
Ok(all_models)
|
||||
} else {
|
||||
Ok(recommended_models)
|
||||
Ok(inventory_models)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -9,7 +9,9 @@ use crate::acp::{
|
||||
use crate::config::search_path::SearchPaths;
|
||||
use crate::config::{Config, GooseMode};
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity};
|
||||
use crate::providers::base::{ProviderDef, ProviderMetadata};
|
||||
use crate::providers::inventory::InventoryIdentityInput;
|
||||
|
||||
const CLAUDE_ACP_PROVIDER_NAME: &str = "claude-acp";
|
||||
const CLAUDE_ACP_DOC_URL: &str = "https://github.com/zed-industries/claude-agent-acp";
|
||||
@@ -78,4 +80,16 @@ impl ProviderDef for ClaudeAcpProvider {
|
||||
AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await
|
||||
})
|
||||
}
|
||||
|
||||
fn supports_inventory_refresh() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn inventory_identity() -> Result<InventoryIdentityInput> {
|
||||
acp_inventory_identity(CLAUDE_ACP_PROVIDER_NAME, CLAUDE_ACP_BINARY)
|
||||
}
|
||||
|
||||
fn inventory_configured() -> bool {
|
||||
acp_adapter_installed(CLAUDE_ACP_BINARY)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,7 +9,9 @@ use crate::acp::{
|
||||
use crate::config::search_path::SearchPaths;
|
||||
use crate::config::{Config, GooseMode};
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity};
|
||||
use crate::providers::base::{ProviderDef, ProviderMetadata};
|
||||
use crate::providers::inventory::InventoryIdentityInput;
|
||||
|
||||
const CODEX_ACP_PROVIDER_NAME: &str = "codex-acp";
|
||||
const CODEX_ACP_DOC_URL: &str = "https://github.com/zed-industries/codex-acp";
|
||||
@@ -98,6 +100,18 @@ impl ProviderDef for CodexAcpProvider {
|
||||
AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await
|
||||
})
|
||||
}
|
||||
|
||||
fn supports_inventory_refresh() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn inventory_identity() -> Result<InventoryIdentityInput> {
|
||||
acp_inventory_identity(CODEX_ACP_PROVIDER_NAME, CODEX_ACP_PROVIDER_NAME)
|
||||
}
|
||||
|
||||
fn inventory_configured() -> bool {
|
||||
acp_adapter_installed(CODEX_ACP_PROVIDER_NAME)
|
||||
}
|
||||
}
|
||||
|
||||
// Codex sandbox scope determines what needs approval: operations within the
|
||||
|
||||
@@ -9,7 +9,9 @@ use crate::acp::{
|
||||
use crate::config::search_path::SearchPaths;
|
||||
use crate::config::{Config, GooseMode};
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity};
|
||||
use crate::providers::base::{ProviderDef, ProviderMetadata};
|
||||
use crate::providers::inventory::InventoryIdentityInput;
|
||||
|
||||
const COPILOT_ACP_PROVIDER_NAME: &str = "copilot-acp";
|
||||
const COPILOT_ACP_DOC_URL: &str = "https://github.com/github/copilot-cli";
|
||||
@@ -84,4 +86,16 @@ impl ProviderDef for CopilotAcpProvider {
|
||||
AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await
|
||||
})
|
||||
}
|
||||
|
||||
fn supports_inventory_refresh() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn inventory_identity() -> Result<InventoryIdentityInput> {
|
||||
acp_inventory_identity(COPILOT_ACP_PROVIDER_NAME, COPILOT_ACP_BINARY)
|
||||
}
|
||||
|
||||
fn inventory_configured() -> bool {
|
||||
acp_adapter_installed(COPILOT_ACP_BINARY)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -340,6 +340,10 @@ impl ProviderDef for DatabricksProvider {
|
||||
) -> BoxFuture<'static, Result<Self::Provider>> {
|
||||
Box::pin(Self::from_env(model))
|
||||
}
|
||||
|
||||
fn supports_inventory_refresh() -> bool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
|
||||
@@ -1100,12 +1100,16 @@ mod tests {
|
||||
fn test_create_request_enabled_thinking_with_budget() -> anyhow::Result<()> {
|
||||
let _guard = env_lock::lock_env([
|
||||
("CLAUDE_THINKING_TYPE", None::<&str>),
|
||||
("CLAUDE_THINKING_ENABLED", Some("1")),
|
||||
("CLAUDE_THINKING_ENABLED", None::<&str>),
|
||||
("CLAUDE_THINKING_BUDGET", Some("10000")),
|
||||
]);
|
||||
|
||||
let mut model_config = ModelConfig::new_or_fail("databricks-claude-3-7-sonnet");
|
||||
model_config.max_tokens = Some(4096);
|
||||
model_config = model_config.with_request_params(Some(std::collections::HashMap::from([(
|
||||
"thinking_type".to_string(),
|
||||
json!("enabled"),
|
||||
)])));
|
||||
|
||||
let request = create_request(&model_config, "system", &[], &[], &ImageFormat::OpenAi)?;
|
||||
|
||||
|
||||
@@ -149,6 +149,10 @@ pub async fn get_from_registry(name: &str) -> Result<ProviderEntry> {
|
||||
.cloned()
|
||||
}
|
||||
|
||||
pub async fn inventory_identity(name: &str) -> Result<super::inventory::InventoryIdentityInput> {
|
||||
get_from_registry(name).await?.inventory_identity()
|
||||
}
|
||||
|
||||
pub async fn create(
|
||||
name: &str,
|
||||
model: ModelConfig,
|
||||
|
||||
@@ -0,0 +1,970 @@
|
||||
use super::base::{ConfigKey, ModelInfo};
|
||||
use super::canonical::{map_provider_name, map_to_canonical_model, CanonicalModelRegistry};
|
||||
use crate::config::declarative_providers::{DeclarativeProviderConfig, ProviderEngine};
|
||||
use crate::config::Config;
|
||||
use crate::session::session_manager::SessionStorage;
|
||||
use anyhow::Result;
|
||||
use chrono::{DateTime, Duration, Utc};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
use sqlx::{Pool, Row, Sqlite, Transaction};
|
||||
use std::collections::{BTreeMap, HashMap, HashSet};
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
const STALE_AFTER_HOURS: i64 = 24;
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct ProviderInventoryEntry {
|
||||
pub provider_id: String,
|
||||
pub provider_name: String,
|
||||
pub configured: bool,
|
||||
pub supports_refresh: bool,
|
||||
pub refreshing: bool,
|
||||
pub models: Vec<InventoryModel>,
|
||||
pub last_updated_at: Option<DateTime<Utc>>,
|
||||
pub last_refresh_attempt_at: Option<DateTime<Utc>>,
|
||||
pub last_refresh_error: Option<String>,
|
||||
pub model_selection_hint: Option<String>,
|
||||
}
|
||||
|
||||
/// Families whose latest model should be surfaced in the compact picker.
|
||||
/// Each entry is matched against the `family` field of enriched models.
|
||||
const RECOMMENDED_FAMILIES: &[&str] = &[
|
||||
"claude-opus",
|
||||
"claude-sonnet",
|
||||
"gpt",
|
||||
"gpt-mini",
|
||||
"glm",
|
||||
"gemini-pro",
|
||||
"gemini-flash",
|
||||
"gemma",
|
||||
];
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase")]
|
||||
pub struct InventoryModel {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub family: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub context_limit: Option<usize>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning: Option<bool>,
|
||||
/// Whether this model should appear in the compact recommended picker.
|
||||
pub recommended: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct InventoryIdentity {
|
||||
pub provider_id: String,
|
||||
pub provider_family: String,
|
||||
pub inventory_key: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct InventoryIdentityInput {
|
||||
pub provider_id: String,
|
||||
pub provider_family: String,
|
||||
pub public_inputs: BTreeMap<String, String>,
|
||||
pub secret_inputs: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl InventoryIdentityInput {
|
||||
pub fn new(
|
||||
provider_id: impl Into<String>,
|
||||
provider_family: impl Into<String>,
|
||||
) -> InventoryIdentityInput {
|
||||
InventoryIdentityInput {
|
||||
provider_id: provider_id.into(),
|
||||
provider_family: provider_family.into(),
|
||||
public_inputs: BTreeMap::new(),
|
||||
secret_inputs: BTreeMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_public(
|
||||
mut self,
|
||||
key: impl Into<String>,
|
||||
value: impl Into<String>,
|
||||
) -> InventoryIdentityInput {
|
||||
self.public_inputs.insert(key.into(), value.into());
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_secret(
|
||||
mut self,
|
||||
key: impl Into<String>,
|
||||
value: impl Into<String>,
|
||||
) -> InventoryIdentityInput {
|
||||
self.secret_inputs.insert(key.into(), value.into());
|
||||
self
|
||||
}
|
||||
|
||||
pub fn into_identity(self) -> Result<InventoryIdentity> {
|
||||
let InventoryIdentityInput {
|
||||
provider_id,
|
||||
provider_family,
|
||||
public_inputs,
|
||||
secret_inputs,
|
||||
} = self;
|
||||
let payload = serde_json::json!({
|
||||
"provider_family": provider_family,
|
||||
"public_inputs": public_inputs,
|
||||
"secret_inputs": secret_inputs,
|
||||
});
|
||||
let digest = Sha256::digest(serde_json::to_vec(&payload)?);
|
||||
Ok(InventoryIdentity {
|
||||
provider_id,
|
||||
provider_family,
|
||||
inventory_key: format!("{digest:x}"),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum RefreshSkipReason {
|
||||
UnknownProvider,
|
||||
NotConfigured,
|
||||
DoesNotSupportRefresh,
|
||||
AlreadyRefreshing,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct RefreshSkip {
|
||||
pub provider_id: String,
|
||||
pub reason: RefreshSkipReason,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct RefreshPlan {
|
||||
pub started: Vec<String>,
|
||||
pub skipped: Vec<RefreshSkip>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ProviderInventoryService {
|
||||
storage: Arc<SessionStorage>,
|
||||
refreshing_keys: Arc<RwLock<HashSet<String>>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct InventorySnapshot {
|
||||
models: Vec<InventoryModel>,
|
||||
last_updated_at: Option<DateTime<Utc>>,
|
||||
last_refresh_attempt_at: Option<DateTime<Utc>>,
|
||||
last_refresh_error: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct ProviderDescriptor {
|
||||
provider_id: String,
|
||||
provider_name: String,
|
||||
identity: InventoryIdentity,
|
||||
configured: bool,
|
||||
supports_refresh: bool,
|
||||
static_models: Vec<ModelInfo>,
|
||||
model_selection_hint: Option<String>,
|
||||
}
|
||||
|
||||
impl ProviderInventoryService {
|
||||
pub fn new(storage: Arc<SessionStorage>) -> ProviderInventoryService {
|
||||
ProviderInventoryService {
|
||||
storage,
|
||||
refreshing_keys: Arc::new(RwLock::new(HashSet::new())),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn entry_for_provider(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<Option<ProviderInventoryEntry>> {
|
||||
let Some(descriptor) = self.describe_provider(provider_id).await? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let snapshot = self.read_snapshot(&descriptor.identity).await?;
|
||||
let refreshing = self
|
||||
.refreshing_keys
|
||||
.read()
|
||||
.await
|
||||
.contains(&descriptor.identity.inventory_key);
|
||||
let models = inventory_models_from_snapshot(
|
||||
snapshot.as_ref(),
|
||||
&descriptor.identity.provider_family,
|
||||
&descriptor.static_models,
|
||||
);
|
||||
|
||||
Ok(Some(ProviderInventoryEntry {
|
||||
provider_id: descriptor.provider_id,
|
||||
provider_name: descriptor.provider_name,
|
||||
configured: descriptor.configured,
|
||||
supports_refresh: descriptor.supports_refresh,
|
||||
refreshing,
|
||||
models,
|
||||
last_updated_at: snapshot
|
||||
.as_ref()
|
||||
.and_then(|snapshot| snapshot.last_updated_at),
|
||||
last_refresh_attempt_at: snapshot
|
||||
.as_ref()
|
||||
.and_then(|snapshot| snapshot.last_refresh_attempt_at),
|
||||
last_refresh_error: snapshot.and_then(|snapshot| snapshot.last_refresh_error),
|
||||
model_selection_hint: descriptor.model_selection_hint,
|
||||
}))
|
||||
}
|
||||
|
||||
pub async fn entries(&self, provider_ids: &[String]) -> Result<Vec<ProviderInventoryEntry>> {
|
||||
let ids = self.resolve_provider_ids(provider_ids).await;
|
||||
let mut entries = Vec::with_capacity(ids.len());
|
||||
for provider_id in ids {
|
||||
if let Some(entry) = self.entry_for_provider(&provider_id).await? {
|
||||
entries.push(entry);
|
||||
}
|
||||
}
|
||||
Ok(entries)
|
||||
}
|
||||
|
||||
pub async fn plan_refresh(&self, provider_ids: &[String]) -> Result<RefreshPlan> {
|
||||
let ids = self.resolve_provider_ids(provider_ids).await;
|
||||
let mut plan = RefreshPlan::default();
|
||||
|
||||
for provider_id in ids {
|
||||
let Some(descriptor) = self.describe_provider(&provider_id).await? else {
|
||||
plan.skipped.push(RefreshSkip {
|
||||
provider_id,
|
||||
reason: RefreshSkipReason::UnknownProvider,
|
||||
});
|
||||
continue;
|
||||
};
|
||||
|
||||
if !descriptor.supports_refresh {
|
||||
plan.skipped.push(RefreshSkip {
|
||||
provider_id: descriptor.provider_id,
|
||||
reason: RefreshSkipReason::DoesNotSupportRefresh,
|
||||
});
|
||||
continue;
|
||||
}
|
||||
|
||||
if !descriptor.configured {
|
||||
plan.skipped.push(RefreshSkip {
|
||||
provider_id: descriptor.provider_id,
|
||||
reason: RefreshSkipReason::NotConfigured,
|
||||
});
|
||||
continue;
|
||||
}
|
||||
|
||||
let mut refreshing_keys = self.refreshing_keys.write().await;
|
||||
if refreshing_keys.contains(&descriptor.identity.inventory_key) {
|
||||
plan.skipped.push(RefreshSkip {
|
||||
provider_id: descriptor.provider_id,
|
||||
reason: RefreshSkipReason::AlreadyRefreshing,
|
||||
});
|
||||
continue;
|
||||
}
|
||||
|
||||
refreshing_keys.insert(descriptor.identity.inventory_key.clone());
|
||||
drop(refreshing_keys);
|
||||
|
||||
self.mark_refresh_started(&descriptor.identity).await?;
|
||||
plan.started.push(descriptor.provider_id);
|
||||
}
|
||||
|
||||
Ok(plan)
|
||||
}
|
||||
|
||||
pub async fn store_refreshed_models(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
model_ids: &[String],
|
||||
) -> Result<()> {
|
||||
let descriptor = self.require_provider(provider_id).await?;
|
||||
let models =
|
||||
enrich_model_ids_with_canonical(&descriptor.identity.provider_family, model_ids);
|
||||
let now = Utc::now();
|
||||
let pool = self.storage.pool().await?;
|
||||
let mut tx = pool.begin().await?;
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO provider_inventory_entries (
|
||||
inventory_key,
|
||||
provider_id,
|
||||
provider_family,
|
||||
last_updated_at,
|
||||
last_refresh_attempt_at,
|
||||
last_refresh_error,
|
||||
updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, NULL, CURRENT_TIMESTAMP)
|
||||
ON CONFLICT(inventory_key) DO UPDATE SET
|
||||
provider_id = excluded.provider_id,
|
||||
provider_family = excluded.provider_family,
|
||||
last_updated_at = excluded.last_updated_at,
|
||||
last_refresh_attempt_at = excluded.last_refresh_attempt_at,
|
||||
last_refresh_error = NULL,
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
"#,
|
||||
)
|
||||
.bind(&descriptor.identity.inventory_key)
|
||||
.bind(&descriptor.identity.provider_id)
|
||||
.bind(&descriptor.identity.provider_family)
|
||||
.bind(now.to_rfc3339())
|
||||
.bind(now.to_rfc3339())
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
|
||||
sqlx::query("DELETE FROM provider_inventory_models WHERE inventory_key = ?")
|
||||
.bind(&descriptor.identity.inventory_key)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
|
||||
for (ordinal, model) in models.iter().enumerate() {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO provider_inventory_models (
|
||||
inventory_key,
|
||||
ordinal,
|
||||
model_id,
|
||||
name,
|
||||
family,
|
||||
context_limit,
|
||||
reasoning,
|
||||
recommended
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&descriptor.identity.inventory_key)
|
||||
.bind(i64::try_from(ordinal)?)
|
||||
.bind(&model.id)
|
||||
.bind(&model.name)
|
||||
.bind(&model.family)
|
||||
.bind(model.context_limit.map(i64::try_from).transpose()?)
|
||||
.bind(model.reasoning)
|
||||
.bind(model.recommended)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
}
|
||||
|
||||
tx.commit().await?;
|
||||
self.refreshing_keys
|
||||
.write()
|
||||
.await
|
||||
.remove(&descriptor.identity.inventory_key);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn store_refresh_error(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
error: impl Into<String>,
|
||||
) -> Result<()> {
|
||||
let descriptor = self.require_provider(provider_id).await?;
|
||||
let error = error.into();
|
||||
let existing = self.read_snapshot(&descriptor.identity).await?;
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO provider_inventory_entries (
|
||||
inventory_key,
|
||||
provider_id,
|
||||
provider_family,
|
||||
last_updated_at,
|
||||
last_refresh_attempt_at,
|
||||
last_refresh_error,
|
||||
updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP)
|
||||
ON CONFLICT(inventory_key) DO UPDATE SET
|
||||
provider_id = excluded.provider_id,
|
||||
provider_family = excluded.provider_family,
|
||||
last_updated_at = excluded.last_updated_at,
|
||||
last_refresh_attempt_at = excluded.last_refresh_attempt_at,
|
||||
last_refresh_error = excluded.last_refresh_error,
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
"#,
|
||||
)
|
||||
.bind(&descriptor.identity.inventory_key)
|
||||
.bind(&descriptor.identity.provider_id)
|
||||
.bind(&descriptor.identity.provider_family)
|
||||
.bind(existing.and_then(|snapshot| snapshot.last_updated_at.map(|time| time.to_rfc3339())))
|
||||
.bind(Utc::now().to_rfc3339())
|
||||
.bind(error)
|
||||
.execute(self.storage.pool().await?)
|
||||
.await?;
|
||||
|
||||
self.refreshing_keys
|
||||
.write()
|
||||
.await
|
||||
.remove(&descriptor.identity.inventory_key);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn is_stale(entry: &ProviderInventoryEntry) -> bool {
|
||||
let Some(last_updated_at) = entry.last_updated_at else {
|
||||
return false;
|
||||
};
|
||||
entry.supports_refresh && Utc::now() - last_updated_at > Duration::hours(STALE_AFTER_HOURS)
|
||||
}
|
||||
|
||||
async fn describe_provider(&self, provider_id: &str) -> Result<Option<ProviderDescriptor>> {
|
||||
let entry = match crate::providers::get_from_registry(provider_id).await {
|
||||
Ok(entry) => entry,
|
||||
Err(_) => return Ok(None),
|
||||
};
|
||||
let metadata = entry.metadata().clone();
|
||||
let identity = crate::providers::inventory_identity(provider_id)
|
||||
.await
|
||||
.unwrap_or_else(|_| fallback_inventory_identity(provider_id))
|
||||
.into_identity()?;
|
||||
|
||||
Ok(Some(ProviderDescriptor {
|
||||
provider_id: metadata.name.clone(),
|
||||
provider_name: metadata.display_name.clone(),
|
||||
identity,
|
||||
configured: entry.inventory_configured(),
|
||||
supports_refresh: entry.supports_inventory_refresh(),
|
||||
static_models: metadata.known_models,
|
||||
model_selection_hint: metadata.model_selection_hint,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn require_provider(&self, provider_id: &str) -> Result<ProviderDescriptor> {
|
||||
self.describe_provider(provider_id)
|
||||
.await?
|
||||
.ok_or_else(|| anyhow::anyhow!("Unknown provider: {}", provider_id))
|
||||
}
|
||||
|
||||
async fn mark_refresh_started(&self, identity: &InventoryIdentity) -> Result<()> {
|
||||
let existing = self.read_snapshot(identity).await?;
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO provider_inventory_entries (
|
||||
inventory_key,
|
||||
provider_id,
|
||||
provider_family,
|
||||
last_updated_at,
|
||||
last_refresh_attempt_at,
|
||||
last_refresh_error,
|
||||
updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, NULL, CURRENT_TIMESTAMP)
|
||||
ON CONFLICT(inventory_key) DO UPDATE SET
|
||||
provider_id = excluded.provider_id,
|
||||
provider_family = excluded.provider_family,
|
||||
last_updated_at = excluded.last_updated_at,
|
||||
last_refresh_attempt_at = excluded.last_refresh_attempt_at,
|
||||
last_refresh_error = NULL,
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
"#,
|
||||
)
|
||||
.bind(&identity.inventory_key)
|
||||
.bind(&identity.provider_id)
|
||||
.bind(&identity.provider_family)
|
||||
.bind(existing.and_then(|snapshot| snapshot.last_updated_at.map(|time| time.to_rfc3339())))
|
||||
.bind(Utc::now().to_rfc3339())
|
||||
.execute(self.storage.pool().await?)
|
||||
.await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn read_snapshot(
|
||||
&self,
|
||||
identity: &InventoryIdentity,
|
||||
) -> Result<Option<InventorySnapshot>> {
|
||||
let pool = self.storage.pool().await?;
|
||||
let entry = sqlx::query(
|
||||
r#"
|
||||
SELECT last_updated_at, last_refresh_attempt_at, last_refresh_error
|
||||
FROM provider_inventory_entries
|
||||
WHERE inventory_key = ?
|
||||
"#,
|
||||
)
|
||||
.bind(&identity.inventory_key)
|
||||
.fetch_optional(pool)
|
||||
.await?;
|
||||
|
||||
let Some(entry) = entry else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let last_updated_at = parse_optional_datetime(entry.try_get("last_updated_at")?)?;
|
||||
let last_refresh_attempt_at =
|
||||
parse_optional_datetime(entry.try_get("last_refresh_attempt_at")?)?;
|
||||
let last_refresh_error = entry.try_get("last_refresh_error")?;
|
||||
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT model_id, name, family, context_limit, reasoning, recommended
|
||||
FROM provider_inventory_models
|
||||
WHERE inventory_key = ?
|
||||
ORDER BY ordinal
|
||||
"#,
|
||||
)
|
||||
.bind(&identity.inventory_key)
|
||||
.fetch_all(pool)
|
||||
.await?;
|
||||
|
||||
let models = rows
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
Ok(InventoryModel {
|
||||
id: row.try_get("model_id")?,
|
||||
name: row.try_get("name")?,
|
||||
family: row.try_get("family")?,
|
||||
context_limit: row
|
||||
.try_get::<Option<i64>, _>("context_limit")?
|
||||
.map(usize::try_from)
|
||||
.transpose()?,
|
||||
reasoning: row.try_get("reasoning")?,
|
||||
recommended: row
|
||||
.try_get::<Option<bool>, _>("recommended")?
|
||||
.unwrap_or(false),
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>, anyhow::Error>>()?;
|
||||
|
||||
Ok(Some(InventorySnapshot {
|
||||
models,
|
||||
last_updated_at,
|
||||
last_refresh_attempt_at,
|
||||
last_refresh_error,
|
||||
}))
|
||||
}
|
||||
|
||||
async fn resolve_provider_ids(&self, provider_ids: &[String]) -> Vec<String> {
|
||||
let mut ids = if provider_ids.is_empty() {
|
||||
crate::providers::providers()
|
||||
.await
|
||||
.into_iter()
|
||||
.map(|(metadata, _)| metadata.name)
|
||||
.collect::<Vec<_>>()
|
||||
} else {
|
||||
provider_ids.to_vec()
|
||||
};
|
||||
ids.sort();
|
||||
ids.dedup();
|
||||
ids
|
||||
}
|
||||
}
|
||||
|
||||
pub fn default_inventory_identity(
|
||||
provider_id: &str,
|
||||
provider_family: &str,
|
||||
config_keys: &[ConfigKey],
|
||||
config: &Config,
|
||||
) -> InventoryIdentityInput {
|
||||
let mut identity = InventoryIdentityInput::new(provider_id, provider_family);
|
||||
|
||||
for key in config_keys {
|
||||
if key.secret {
|
||||
if let Some(value) = config_secret_value(config, &key.name) {
|
||||
identity.secret_inputs.insert(key.name.clone(), value);
|
||||
}
|
||||
} else if let Some(value) = config_param_value(config, &key.name) {
|
||||
identity.public_inputs.insert(key.name.clone(), value);
|
||||
}
|
||||
}
|
||||
|
||||
identity
|
||||
}
|
||||
|
||||
pub fn default_inventory_configured(config_keys: &[ConfigKey], config: &Config) -> bool {
|
||||
config_keys.iter().all(|key| {
|
||||
if !key.required {
|
||||
return true;
|
||||
}
|
||||
if key.default.is_some() {
|
||||
return true;
|
||||
}
|
||||
if key.secret {
|
||||
config.get_secret::<serde_json::Value>(&key.name).is_ok()
|
||||
} else {
|
||||
config.get_param::<serde_json::Value>(&key.name).is_ok()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
pub fn declarative_inventory_identity(
|
||||
config: &DeclarativeProviderConfig,
|
||||
) -> Result<InventoryIdentityInput> {
|
||||
let global = Config::global();
|
||||
let mut identity = InventoryIdentityInput::new(
|
||||
config.name.clone(),
|
||||
config
|
||||
.catalog_provider_id
|
||||
.clone()
|
||||
.unwrap_or_else(|| match config.engine {
|
||||
ProviderEngine::OpenAI => "openai".to_string(),
|
||||
ProviderEngine::Anthropic => "anthropic".to_string(),
|
||||
ProviderEngine::Ollama => "ollama".to_string(),
|
||||
}),
|
||||
);
|
||||
|
||||
identity
|
||||
.public_inputs
|
||||
.insert("base_url".to_string(), config.base_url.clone());
|
||||
|
||||
if let Some(base_path) = &config.base_path {
|
||||
identity
|
||||
.public_inputs
|
||||
.insert("base_path".to_string(), base_path.clone());
|
||||
}
|
||||
if let Some(catalog_provider_id) = &config.catalog_provider_id {
|
||||
identity.public_inputs.insert(
|
||||
"catalog_provider_id".to_string(),
|
||||
catalog_provider_id.clone(),
|
||||
);
|
||||
}
|
||||
if let Some(dynamic_models) = config.dynamic_models {
|
||||
identity
|
||||
.public_inputs
|
||||
.insert("dynamic_models".to_string(), dynamic_models.to_string());
|
||||
}
|
||||
identity.public_inputs.insert(
|
||||
"skip_canonical_filtering".to_string(),
|
||||
config.skip_canonical_filtering.to_string(),
|
||||
);
|
||||
if !config.models.is_empty() {
|
||||
identity.public_inputs.insert(
|
||||
"models".to_string(),
|
||||
serde_json::to_string(
|
||||
&config
|
||||
.models
|
||||
.iter()
|
||||
.map(|model| &model.name)
|
||||
.collect::<Vec<_>>(),
|
||||
)?,
|
||||
);
|
||||
}
|
||||
if let Some(headers) = &config.headers {
|
||||
identity
|
||||
.public_inputs
|
||||
.insert("headers".to_string(), serialize_string_map(headers)?);
|
||||
}
|
||||
if config.requires_auth && !config.api_key_env.is_empty() {
|
||||
if let Some(value) = config_secret_value(global, &config.api_key_env) {
|
||||
identity
|
||||
.secret_inputs
|
||||
.insert(config.api_key_env.clone(), value);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(identity)
|
||||
}
|
||||
|
||||
pub fn config_param_value(config: &Config, key: &str) -> Option<String> {
|
||||
config
|
||||
.get_param::<serde_json::Value>(key)
|
||||
.ok()
|
||||
.and_then(|value| normalize_json_value(&value))
|
||||
}
|
||||
|
||||
pub fn config_secret_value(config: &Config, key: &str) -> Option<String> {
|
||||
config
|
||||
.get_secret::<serde_json::Value>(key)
|
||||
.ok()
|
||||
.and_then(|value| normalize_json_value(&value))
|
||||
}
|
||||
|
||||
pub fn serialize_string_map(map: &HashMap<String, String>) -> Result<String> {
|
||||
let ordered = map
|
||||
.iter()
|
||||
.map(|(key, value)| (key.clone(), value.clone()))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
Ok(serde_json::to_string(&ordered)?)
|
||||
}
|
||||
|
||||
fn parse_optional_datetime(value: Option<String>) -> Result<Option<DateTime<Utc>>> {
|
||||
value
|
||||
.map(|value| value.parse::<DateTime<Utc>>())
|
||||
.transpose()
|
||||
.map_err(Into::into)
|
||||
}
|
||||
|
||||
fn normalize_json_value(value: &serde_json::Value) -> Option<String> {
|
||||
match value {
|
||||
serde_json::Value::Null => None,
|
||||
serde_json::Value::String(value) if value.is_empty() => None,
|
||||
serde_json::Value::String(value) => Some(value.clone()),
|
||||
other => serde_json::to_string(other).ok(),
|
||||
}
|
||||
}
|
||||
|
||||
fn fallback_inventory_identity(provider_id: &str) -> InventoryIdentityInput {
|
||||
InventoryIdentityInput::new(
|
||||
provider_id.to_string(),
|
||||
map_provider_name(provider_id).to_string(),
|
||||
)
|
||||
}
|
||||
|
||||
fn enrich_model_ids_with_canonical(
|
||||
provider_family: &str,
|
||||
model_ids: &[String],
|
||||
) -> Vec<InventoryModel> {
|
||||
let mut models: Vec<InventoryModel> = Vec::new();
|
||||
let mut seen_names: HashSet<String> = HashSet::new();
|
||||
|
||||
for id in model_ids {
|
||||
let model = enriched_model(provider_family, id, None);
|
||||
if !seen_names.insert(model.name.clone()) {
|
||||
continue;
|
||||
}
|
||||
models.push(model);
|
||||
}
|
||||
|
||||
// For databricks, prefer goose- prefixed model_ids when there are duplicates.
|
||||
// Re-scan: if a later model_id with "goose-" prefix maps to the same display name,
|
||||
// swap it in.
|
||||
if provider_family == "databricks" {
|
||||
let mut name_to_idx: HashMap<String, usize> = HashMap::new();
|
||||
for (idx, model) in models.iter().enumerate() {
|
||||
name_to_idx.insert(model.name.clone(), idx);
|
||||
}
|
||||
for id in model_ids {
|
||||
if !id.starts_with("goose-") {
|
||||
continue;
|
||||
}
|
||||
let candidate = enriched_model(provider_family, id, None);
|
||||
if let Some(&idx) = name_to_idx.get(&candidate.name) {
|
||||
if !models[idx].id.starts_with("goose-") {
|
||||
models[idx].id = candidate.id;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Mark the latest model per recommended family.
|
||||
let mut seen_recommended_families: HashSet<String> = HashSet::new();
|
||||
for model in &mut models {
|
||||
if let Some(family) = &model.family {
|
||||
if RECOMMENDED_FAMILIES.contains(&family.as_str())
|
||||
&& seen_recommended_families.insert(family.clone())
|
||||
{
|
||||
model.recommended = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
models
|
||||
}
|
||||
|
||||
fn configured_models_to_inventory(
|
||||
provider_family: &str,
|
||||
models: &[ModelInfo],
|
||||
) -> Vec<InventoryModel> {
|
||||
let mut result: Vec<InventoryModel> = Vec::new();
|
||||
let mut seen_names: HashSet<String> = HashSet::new();
|
||||
for model in models {
|
||||
let enriched = enriched_model(provider_family, &model.name, Some(model.context_limit));
|
||||
if seen_names.insert(enriched.name.clone()) {
|
||||
result.push(enriched);
|
||||
}
|
||||
}
|
||||
|
||||
let mut seen_recommended_families: HashSet<String> = HashSet::new();
|
||||
for model in &mut result {
|
||||
if let Some(family) = &model.family {
|
||||
if RECOMMENDED_FAMILIES.contains(&family.as_str())
|
||||
&& seen_recommended_families.insert(family.clone())
|
||||
{
|
||||
model.recommended = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
fn inventory_models_from_snapshot(
|
||||
snapshot: Option<&InventorySnapshot>,
|
||||
provider_family: &str,
|
||||
configured_models: &[ModelInfo],
|
||||
) -> Vec<InventoryModel> {
|
||||
match snapshot {
|
||||
Some(snapshot) if !snapshot.models.is_empty() || snapshot.last_updated_at.is_some() => {
|
||||
snapshot.models.clone()
|
||||
}
|
||||
_ => configured_models_to_inventory(provider_family, configured_models),
|
||||
}
|
||||
}
|
||||
|
||||
fn enriched_model(
|
||||
provider_family: &str,
|
||||
model_id: &str,
|
||||
fallback_context_limit: Option<usize>,
|
||||
) -> InventoryModel {
|
||||
let registry = CanonicalModelRegistry::bundled().ok();
|
||||
let canonical = registry.as_ref().and_then(|registry| {
|
||||
let canonical_id = map_to_canonical_model(provider_family, model_id, registry)?;
|
||||
let (provider, model) = canonical_id.split_once('/')?;
|
||||
registry.get(provider, model).cloned()
|
||||
});
|
||||
|
||||
InventoryModel {
|
||||
id: model_id.to_string(),
|
||||
name: canonical
|
||||
.as_ref()
|
||||
.map(|model| model.name.clone())
|
||||
.unwrap_or_else(|| model_id.to_string()),
|
||||
family: canonical.as_ref().and_then(|model| model.family.clone()),
|
||||
context_limit: canonical
|
||||
.as_ref()
|
||||
.map(|model| model.limit.context)
|
||||
.or(fallback_context_limit),
|
||||
reasoning: canonical.as_ref().and_then(|model| model.reasoning),
|
||||
recommended: false,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn create_tables(pool: &Pool<Sqlite>) -> Result<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
CREATE TABLE IF NOT EXISTS provider_inventory_entries (
|
||||
inventory_key TEXT PRIMARY KEY,
|
||||
provider_id TEXT NOT NULL,
|
||||
provider_family TEXT NOT NULL,
|
||||
last_updated_at TEXT,
|
||||
last_refresh_attempt_at TEXT,
|
||||
last_refresh_error TEXT,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await?;
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
CREATE TABLE IF NOT EXISTS provider_inventory_models (
|
||||
inventory_key TEXT NOT NULL REFERENCES provider_inventory_entries(inventory_key) ON DELETE CASCADE,
|
||||
ordinal INTEGER NOT NULL,
|
||||
model_id TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
family TEXT,
|
||||
context_limit INTEGER,
|
||||
reasoning BOOLEAN,
|
||||
recommended BOOLEAN,
|
||||
PRIMARY KEY (inventory_key, ordinal)
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await?;
|
||||
|
||||
sqlx::query(
|
||||
"CREATE INDEX IF NOT EXISTS idx_provider_inventory_provider_id ON provider_inventory_entries(provider_id)",
|
||||
)
|
||||
.execute(pool)
|
||||
.await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn create_tables_in_tx(tx: &mut Transaction<'_, Sqlite>) -> Result<()> {
|
||||
sqlx::query(
|
||||
r#"
|
||||
CREATE TABLE IF NOT EXISTS provider_inventory_entries (
|
||||
inventory_key TEXT PRIMARY KEY,
|
||||
provider_id TEXT NOT NULL,
|
||||
provider_family TEXT NOT NULL,
|
||||
last_updated_at TEXT,
|
||||
last_refresh_attempt_at TEXT,
|
||||
last_refresh_error TEXT,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
CREATE TABLE IF NOT EXISTS provider_inventory_models (
|
||||
inventory_key TEXT NOT NULL REFERENCES provider_inventory_entries(inventory_key) ON DELETE CASCADE,
|
||||
ordinal INTEGER NOT NULL,
|
||||
model_id TEXT NOT NULL,
|
||||
name TEXT NOT NULL,
|
||||
family TEXT,
|
||||
context_limit INTEGER,
|
||||
reasoning BOOLEAN,
|
||||
recommended BOOLEAN,
|
||||
PRIMARY KEY (inventory_key, ordinal)
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
|
||||
sqlx::query(
|
||||
"CREATE INDEX IF NOT EXISTS idx_provider_inventory_provider_id ON provider_inventory_entries(provider_id)",
|
||||
)
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn inventory_identity_hash_changes_with_secret_inputs() {
|
||||
let left = InventoryIdentityInput::new("openai", "openai")
|
||||
.with_public("host", "https://api.openai.com")
|
||||
.with_secret("api_key", "secret-a")
|
||||
.into_identity()
|
||||
.unwrap();
|
||||
let right = InventoryIdentityInput::new("openai", "openai")
|
||||
.with_public("host", "https://api.openai.com")
|
||||
.with_secret("api_key", "secret-b")
|
||||
.into_identity()
|
||||
.unwrap();
|
||||
|
||||
assert_ne!(left.inventory_key, right.inventory_key);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn configured_models_use_canonical_enrichment() {
|
||||
let models =
|
||||
configured_models_to_inventory("anthropic", &[ModelInfo::new("claude-sonnet-4-5", 0)]);
|
||||
|
||||
assert_eq!(models.len(), 1);
|
||||
assert!(models[0].name.contains("Claude"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn inventory_uses_configured_models_before_first_successful_refresh() {
|
||||
let configured_models = [ModelInfo::new("claude-sonnet-4-5", 0)];
|
||||
let snapshot = InventorySnapshot {
|
||||
models: vec![],
|
||||
last_updated_at: None,
|
||||
last_refresh_attempt_at: Some(Utc::now()),
|
||||
last_refresh_error: Some("auth failed".to_string()),
|
||||
};
|
||||
|
||||
let models =
|
||||
inventory_models_from_snapshot(Some(&snapshot), "anthropic", &configured_models);
|
||||
|
||||
assert_eq!(models.len(), 1);
|
||||
assert_eq!(models[0].id, "claude-sonnet-4-5");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn inventory_preserves_empty_models_after_successful_refresh() {
|
||||
let configured_models = [ModelInfo::new("claude-sonnet-4-5", 0)];
|
||||
let snapshot = InventorySnapshot {
|
||||
models: vec![],
|
||||
last_updated_at: Some(Utc::now()),
|
||||
last_refresh_attempt_at: Some(Utc::now()),
|
||||
last_refresh_error: None,
|
||||
};
|
||||
|
||||
let models =
|
||||
inventory_models_from_snapshot(Some(&snapshot), "anthropic", &configured_models);
|
||||
|
||||
assert!(models.is_empty());
|
||||
}
|
||||
}
|
||||
@@ -1,3 +1,4 @@
|
||||
mod acp_tooling;
|
||||
pub mod amp_acp;
|
||||
pub mod anthropic;
|
||||
pub mod api_client;
|
||||
@@ -28,6 +29,7 @@ pub mod gemini_oauth;
|
||||
pub mod githubcopilot;
|
||||
pub mod google;
|
||||
mod init;
|
||||
pub mod inventory;
|
||||
pub mod kimicode;
|
||||
pub mod litellm;
|
||||
#[cfg(feature = "local-inference")]
|
||||
@@ -56,6 +58,6 @@ pub mod xai;
|
||||
|
||||
pub use init::{
|
||||
cleanup_provider, create, create_with_default_model, create_with_named_model,
|
||||
get_from_registry, providers, refresh_custom_providers,
|
||||
get_from_registry, inventory_identity, providers, refresh_custom_providers,
|
||||
};
|
||||
pub use retry::{retry_operation, RetryConfig};
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use super::api_client::{ApiClient, AuthMethod};
|
||||
use super::base::{ConfigKey, MessageStream, Provider, ProviderDef, ProviderMetadata};
|
||||
use super::errors::ProviderError;
|
||||
use super::inventory::InventoryIdentityInput;
|
||||
use super::openai_compatible::handle_status_openai_compat;
|
||||
use super::retry::{ProviderRetry, RetryConfig};
|
||||
use super::utils::{ImageFormat, RequestLog};
|
||||
@@ -256,6 +257,22 @@ impl ProviderDef for OllamaProvider {
|
||||
) -> BoxFuture<'static, Result<Self::Provider>> {
|
||||
Box::pin(Self::from_env(model))
|
||||
}
|
||||
|
||||
fn supports_inventory_refresh() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn inventory_identity() -> Result<InventoryIdentityInput> {
|
||||
let config = crate::config::Config::global();
|
||||
Ok(
|
||||
InventoryIdentityInput::new(OLLAMA_PROVIDER_NAME, OLLAMA_PROVIDER_NAME).with_public(
|
||||
"host",
|
||||
config
|
||||
.get_param::<String>("OLLAMA_HOST")
|
||||
.unwrap_or_else(|_| OLLAMA_HOST.to_string()),
|
||||
),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
|
||||
@@ -7,6 +7,7 @@ use super::formats::openai_responses::{
|
||||
create_responses_request, get_responses_usage, responses_api_to_message,
|
||||
responses_api_to_streaming_message, ResponsesApiResponse,
|
||||
};
|
||||
use super::inventory::{config_secret_value, InventoryIdentityInput};
|
||||
use super::openai_compatible::{
|
||||
handle_response_openai_compat, handle_status_openai_compat, stream_openai_compat,
|
||||
};
|
||||
@@ -425,6 +426,58 @@ impl ProviderDef for OpenAiProvider {
|
||||
) -> BoxFuture<'static, Result<Self::Provider>> {
|
||||
Box::pin(Self::from_env(model))
|
||||
}
|
||||
|
||||
fn supports_inventory_refresh() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn inventory_configured() -> bool {
|
||||
let config = crate::config::Config::global();
|
||||
// If the host is explicitly set to something non-default, trust the user's
|
||||
// custom setup (e.g. a local server that doesn't require an API key).
|
||||
if let Ok(host) = config.get_param::<String>("OPENAI_HOST") {
|
||||
if host != "https://api.openai.com" {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
// Standard OpenAI endpoint requires an API key.
|
||||
config
|
||||
.get_secret::<serde_json::Value>("OPENAI_API_KEY")
|
||||
.is_ok()
|
||||
}
|
||||
|
||||
fn inventory_identity() -> Result<InventoryIdentityInput> {
|
||||
let config = crate::config::Config::global();
|
||||
let mut identity =
|
||||
InventoryIdentityInput::new(OPEN_AI_PROVIDER_NAME, OPEN_AI_PROVIDER_NAME)
|
||||
.with_public(
|
||||
"host",
|
||||
config
|
||||
.get_param::<String>("OPENAI_HOST")
|
||||
.unwrap_or_else(|_| "https://api.openai.com".to_string()),
|
||||
)
|
||||
.with_public(
|
||||
"base_path",
|
||||
config
|
||||
.get_param::<String>("OPENAI_BASE_PATH")
|
||||
.unwrap_or_else(|_| OPEN_AI_DEFAULT_BASE_PATH.to_string()),
|
||||
);
|
||||
|
||||
if let Ok(organization) = config.get_param::<String>("OPENAI_ORGANIZATION") {
|
||||
identity = identity.with_public("organization", organization);
|
||||
}
|
||||
if let Ok(project) = config.get_param::<String>("OPENAI_PROJECT") {
|
||||
identity = identity.with_public("project", project);
|
||||
}
|
||||
if let Some(api_key) = config_secret_value(config, "OPENAI_API_KEY") {
|
||||
identity = identity.with_secret("api_key", api_key);
|
||||
}
|
||||
if let Some(custom_headers) = config_secret_value(config, "OPENAI_CUSTOM_HEADERS") {
|
||||
identity = identity.with_secret("custom_headers", custom_headers);
|
||||
}
|
||||
|
||||
Ok(identity)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
|
||||
@@ -9,7 +9,9 @@ use crate::acp::{
|
||||
use crate::config::search_path::SearchPaths;
|
||||
use crate::config::{Config, GooseMode};
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::acp_tooling::{acp_adapter_installed, acp_inventory_identity};
|
||||
use crate::providers::base::{ProviderDef, ProviderMetadata};
|
||||
use crate::providers::inventory::InventoryIdentityInput;
|
||||
|
||||
const PI_ACP_PROVIDER_NAME: &str = "pi-acp";
|
||||
const PI_ACP_DOC_URL: &str = "https://github.com/anthropics/pi";
|
||||
@@ -36,6 +38,7 @@ impl ProviderDef for PiAcpProvider {
|
||||
"Set in your goose config file (`~/.config/goose/config.yaml` on macOS/Linux):\n GOOSE_PROVIDER: pi-acp\n GOOSE_MODEL: current",
|
||||
"Restart goose for changes to take effect",
|
||||
])
|
||||
.with_model_selection_hint("Use the Pi CLI to configure models")
|
||||
}
|
||||
|
||||
fn from_env(
|
||||
@@ -70,4 +73,16 @@ impl ProviderDef for PiAcpProvider {
|
||||
AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await
|
||||
})
|
||||
}
|
||||
|
||||
fn supports_inventory_refresh() -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn inventory_identity() -> Result<InventoryIdentityInput> {
|
||||
acp_inventory_identity(PI_ACP_PROVIDER_NAME, PI_ACP_BINARY)
|
||||
}
|
||||
|
||||
fn inventory_configured() -> bool {
|
||||
acp_adapter_installed(PI_ACP_BINARY)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use super::base::{ModelInfo, Provider, ProviderDef, ProviderMetadata, ProviderType};
|
||||
use super::inventory::InventoryIdentityInput;
|
||||
use crate::config::{DeclarativeProviderConfig, ExtensionConfig};
|
||||
use crate::model::ModelConfig;
|
||||
use anyhow::Result;
|
||||
@@ -14,12 +15,20 @@ pub type ProviderConstructor = Arc<
|
||||
|
||||
pub type ProviderCleanup = Arc<dyn Fn() -> BoxFuture<'static, Result<()>> + Send + Sync>;
|
||||
|
||||
pub type ProviderInventoryIdentityResolver =
|
||||
Arc<dyn Fn() -> Result<InventoryIdentityInput> + Send + Sync>;
|
||||
|
||||
pub type ProviderInventoryConfiguredResolver = Arc<dyn Fn() -> bool + Send + Sync>;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ProviderEntry {
|
||||
metadata: ProviderMetadata,
|
||||
pub(crate) constructor: ProviderConstructor,
|
||||
pub(crate) inventory_identity: ProviderInventoryIdentityResolver,
|
||||
pub(crate) inventory_configured: ProviderInventoryConfiguredResolver,
|
||||
pub(crate) cleanup: Option<ProviderCleanup>,
|
||||
provider_type: ProviderType,
|
||||
supports_inventory_refresh: bool,
|
||||
}
|
||||
|
||||
impl ProviderEntry {
|
||||
@@ -27,6 +36,22 @@ impl ProviderEntry {
|
||||
&self.metadata
|
||||
}
|
||||
|
||||
pub fn provider_type(&self) -> ProviderType {
|
||||
self.provider_type
|
||||
}
|
||||
|
||||
pub fn supports_inventory_refresh(&self) -> bool {
|
||||
self.supports_inventory_refresh
|
||||
}
|
||||
|
||||
pub fn inventory_identity(&self) -> Result<InventoryIdentityInput> {
|
||||
(self.inventory_identity)()
|
||||
}
|
||||
|
||||
pub fn inventory_configured(&self) -> bool {
|
||||
(self.inventory_configured)()
|
||||
}
|
||||
|
||||
fn normalize_model_config(&self, mut model: ModelConfig) -> ModelConfig {
|
||||
model = model.with_canonical_limits(&self.metadata.name);
|
||||
|
||||
@@ -92,24 +117,30 @@ impl ProviderRegistry {
|
||||
Ok(Arc::new(provider) as Arc<dyn Provider>)
|
||||
})
|
||||
}),
|
||||
inventory_identity: Arc::new(F::inventory_identity),
|
||||
inventory_configured: Arc::new(F::inventory_configured),
|
||||
cleanup: None,
|
||||
provider_type: if preferred {
|
||||
ProviderType::Preferred
|
||||
} else {
|
||||
ProviderType::Builtin
|
||||
},
|
||||
supports_inventory_refresh: F::supports_inventory_refresh(),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
pub fn register_with_name<P, F>(
|
||||
pub fn register_with_name<P, F, G>(
|
||||
&mut self,
|
||||
config: &DeclarativeProviderConfig,
|
||||
provider_type: ProviderType,
|
||||
supports_inventory_refresh: bool,
|
||||
constructor: F,
|
||||
inventory_identity: G,
|
||||
) where
|
||||
P: ProviderDef + 'static,
|
||||
F: Fn(ModelConfig) -> Result<P::Provider> + Send + Sync + 'static,
|
||||
G: Fn() -> Result<InventoryIdentityInput> + Send + Sync + 'static,
|
||||
{
|
||||
let base_metadata = P::metadata();
|
||||
let description = config
|
||||
@@ -174,7 +205,9 @@ impl ProviderRegistry {
|
||||
model_doc_link: base_metadata.model_doc_link,
|
||||
config_keys,
|
||||
setup_steps: vec![],
|
||||
model_selection_hint: None,
|
||||
};
|
||||
let inventory_config_keys = custom_metadata.config_keys.clone();
|
||||
|
||||
self.entries.insert(
|
||||
config.name.clone(),
|
||||
@@ -187,8 +220,16 @@ impl ProviderRegistry {
|
||||
Ok(Arc::new(provider) as Arc<dyn Provider>)
|
||||
})
|
||||
}),
|
||||
inventory_identity: Arc::new(inventory_identity),
|
||||
inventory_configured: Arc::new(move || {
|
||||
super::inventory::default_inventory_configured(
|
||||
&inventory_config_keys,
|
||||
crate::config::Config::global(),
|
||||
)
|
||||
}),
|
||||
cleanup: None,
|
||||
provider_type,
|
||||
supports_inventory_refresh,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
@@ -19,7 +19,7 @@ use std::sync::{Arc, LazyLock};
|
||||
use tracing::{info, warn};
|
||||
use utoipa::ToSchema;
|
||||
|
||||
pub const CURRENT_SCHEMA_VERSION: i32 = 10;
|
||||
pub const CURRENT_SCHEMA_VERSION: i32 = 11;
|
||||
pub const SESSIONS_FOLDER: &str = "sessions";
|
||||
pub const DB_NAME: &str = "sessions.db";
|
||||
|
||||
@@ -717,6 +717,8 @@ impl SessionStorage {
|
||||
.execute(pool)
|
||||
.await?;
|
||||
|
||||
crate::providers::inventory::create_tables(pool).await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -1060,6 +1062,9 @@ impl SessionStorage {
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
}
|
||||
11 => {
|
||||
crate::providers::inventory::create_tables_in_tx(tx).await?;
|
||||
}
|
||||
_ => {
|
||||
anyhow::bail!("Unknown migration version: {}", version);
|
||||
}
|
||||
|
||||
@@ -374,6 +374,7 @@ mod tests {
|
||||
model_doc_link: "".to_string(),
|
||||
config_keys: vec![],
|
||||
setup_steps: vec![],
|
||||
model_selection_hint: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -542,6 +543,7 @@ mod tests {
|
||||
model_doc_link: "".to_string(),
|
||||
config_keys: vec![],
|
||||
setup_steps: vec![],
|
||||
model_selection_hint: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -867,6 +869,7 @@ mod tests {
|
||||
model_doc_link: "".to_string(),
|
||||
config_keys: vec![],
|
||||
setup_steps: vec![],
|
||||
model_selection_hint: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -192,6 +192,7 @@ impl ProviderDef for MockCompactionProvider {
|
||||
model_doc_link: "".to_string(),
|
||||
config_keys: vec![],
|
||||
setup_steps: vec![],
|
||||
model_selection_hint: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user