From c5354902ef37788c1dd85f6eb8f613639c95afbe Mon Sep 17 00:00:00 2001 From: Douwe Osinga Date: Tue, 23 Jun 2026 03:29:17 -0400 Subject: [PATCH 01/12] fix(desktop): publish mac updater metadata (#9945) Co-authored-by: Douwe M Osinga --- .github/workflows/release.yml | 8 ++ .../scripts/generate-mac-update-manifest.js | 112 ++++++++++++++++++ .../components/settings/app/UpdateSection.tsx | 32 +++-- ui/desktop/src/i18n/messages/en.json | 8 +- 4 files changed, 144 insertions(+), 16 deletions(-) create mode 100644 ui/desktop/scripts/generate-mac-update-manifest.js diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index c3de49fea59b..35884706aef7 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -108,11 +108,16 @@ jobs: id-token: write # Required for Sigstore OIDC signing attestations: write # Required for SLSA build provenance attestations steps: + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + - name: Download all artifacts uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8.0.1 with: merge-multiple: true + - name: Generate macOS update manifest + run: node ui/desktop/scripts/generate-mac-update-manifest.js --version "${GITHUB_REF_NAME}" --directory . + - name: Attest build provenance uses: actions/attest-build-provenance@a2bbfa25375fe432b6a289bc6b6cd05ecd0c4c32 # v4.1.0 with: @@ -125,6 +130,7 @@ jobs: *.deb *.rpm *.flatpak + latest-mac.yml download_cli.sh # Create/update the versioned release @@ -141,6 +147,7 @@ jobs: *.deb *.rpm *.flatpak + latest-mac.yml download_cli.sh allowUpdates: true omitBody: true @@ -162,6 +169,7 @@ jobs: *.deb *.rpm *.flatpak + latest-mac.yml download_cli.sh allowUpdates: true omitBody: true diff --git a/ui/desktop/scripts/generate-mac-update-manifest.js b/ui/desktop/scripts/generate-mac-update-manifest.js new file mode 100644 index 000000000000..a909802c2162 --- /dev/null +++ b/ui/desktop/scripts/generate-mac-update-manifest.js @@ -0,0 +1,112 @@ +#!/usr/bin/env node + +const crypto = require('node:crypto'); +const fs = require('node:fs'); +const path = require('node:path'); + +function usage() { + console.error( + 'Usage: node scripts/generate-mac-update-manifest.js --version [--directory ]' + ); +} + +function parseArgs(argv) { + const args = { + directory: process.cwd(), + version: '', + }; + + for (let i = 0; i < argv.length; i += 1) { + const arg = argv[i]; + if (arg === '--version') { + args.version = argv[++i] || ''; + } else if (arg === '--directory') { + args.directory = argv[++i] || ''; + } else { + usage(); + process.exit(1); + } + } + + if (!args.version || !args.directory) { + usage(); + process.exit(1); + } + + args.version = args.version.replace(/^v/, ''); + args.directory = path.resolve(args.directory); + return args; +} + +function ensureFile(filePath) { + if (!fs.existsSync(filePath)) { + throw new Error(`Missing required file: ${filePath}`); + } +} + +function copyIfDifferent(source, target) { + ensureFile(source); + if (path.resolve(source) === path.resolve(target)) { + return; + } + fs.copyFileSync(source, target); +} + +function sha512(filePath) { + const hash = crypto.createHash('sha512'); + hash.update(fs.readFileSync(filePath)); + return hash.digest('base64'); +} + +function yamlString(value) { + return JSON.stringify(value); +} + +function writeManifest({ directory, version }) { + const files = [ + { + sourceName: 'Goose.zip', + updateName: 'Goose-darwin-arm64.zip', + }, + { + sourceName: 'Goose_intel_mac.zip', + updateName: 'Goose-darwin-x64.zip', + }, + ]; + + const entries = files.map(({ sourceName, updateName }) => { + const sourcePath = path.join(directory, sourceName); + const updatePath = path.join(directory, updateName); + copyIfDifferent(sourcePath, updatePath); + + const stats = fs.statSync(updatePath); + return { + url: updateName, + sha512: sha512(updatePath), + size: stats.size, + }; + }); + + const manifest = [ + `version: ${yamlString(version)}`, + 'files:', + ...entries.flatMap((entry) => [ + ` - url: ${yamlString(entry.url)}`, + ` sha512: ${yamlString(entry.sha512)}`, + ` size: ${entry.size}`, + ]), + `path: ${yamlString(entries[0].url)}`, + `sha512: ${yamlString(entries[0].sha512)}`, + `releaseDate: ${yamlString(new Date().toISOString())}`, + '', + ].join('\n'); + + fs.writeFileSync(path.join(directory, 'latest-mac.yml'), manifest); +} + +try { + writeManifest(parseArgs(process.argv.slice(2))); +} catch (error) { + console.error(error instanceof Error ? error.message : error); + process.exit(1); +} diff --git a/ui/desktop/src/components/settings/app/UpdateSection.tsx b/ui/desktop/src/components/settings/app/UpdateSection.tsx index d295f1f9d95c..d69045e7c410 100644 --- a/ui/desktop/src/components/settings/app/UpdateSection.tsx +++ b/ui/desktop/src/components/settings/app/UpdateSection.tsx @@ -81,7 +81,8 @@ const i18n = defineMessages({ }, autoDownload: { id: 'updateSection.autoDownload', - defaultMessage: 'Update will be downloaded automatically in the background.', + defaultMessage: + 'Goose will download the update in the background and install it the next time you quit or restart.', }, manualInstallNote: { id: 'updateSection.manualInstallNote', @@ -89,7 +90,7 @@ const i18n = defineMessages({ }, autoInstallNote: { id: 'updateSection.autoInstallNote', - defaultMessage: 'The update will be installed automatically when you quit the app.', + defaultMessage: 'No manual install is needed.', }, readyInstallManual: { id: 'updateSection.readyInstallManual', @@ -101,11 +102,12 @@ const i18n = defineMessages({ }, readyInstallAuto: { id: 'updateSection.readyInstallAuto', - defaultMessage: '✓ Update is ready! It will be installed when you quit Goose.', + defaultMessage: + "✓ Update is ready. Restart Goose to finish installing it, or quit when you're done.", }, installNowHint: { id: 'updateSection.installNowHint', - defaultMessage: 'Or click "Install & Restart" to update now.', + defaultMessage: 'Click "Install & Restart" to update now.', }, }); @@ -176,7 +178,6 @@ export default function UpdateSection() { // Listen for updater events window.electron.onUpdaterEvent((event) => { - switch (event.event) { case 'checking-for-update': setUpdateStatus('checking'); @@ -354,10 +355,15 @@ export default function UpdateSection() {
{updateInfo.currentVersion || intl.formatMessage(i18n.loading)}
-
{intl.formatMessage(i18n.currentVersion)}
+
+ {intl.formatMessage(i18n.currentVersion)} +
{updateInfo.latestVersion && updateInfo.isUpdateAvailable && ( - {intl.formatMessage(i18n.versionAvailable, { version: updateInfo.latestVersion })} + + {' '} + {intl.formatMessage(i18n.versionAvailable, { version: updateInfo.latestVersion })} + )} {updateInfo.currentVersion && updateInfo.isUpdateAvailable === false && ( {intl.formatMessage(i18n.upToDate)} @@ -375,11 +381,13 @@ export default function UpdateSection() { {intl.formatMessage(i18n.checkForUpdates)} - {updateInfo.isUpdateAvailable && updateStatus === 'idle' && autoDownloadEffectivelyDisabled && ( - - )} + {updateInfo.isUpdateAvailable && + updateStatus === 'idle' && + autoDownloadEffectivelyDisabled && ( + + )} {updateStatus === 'ready' && ( + + {extensionType === 'streamable_http' && ( +
+ +
+ {headers.map((header, index) => ( +
+ handleHeaderChange(index, 'name', e.target.value)} + className="flex-1 p-3 border border-borderSubtle rounded-lg bg-bgSubtle text-textStandard" + placeholder="Header Name" + /> + handleHeaderChange(index, 'description', e.target.value)} + className="flex-1 p-3 border border-borderSubtle rounded-lg bg-bgSubtle text-textStandard" + placeholder="Description" + /> + +
+ ))} + +
+
+ )} )} @@ -456,15 +618,16 @@ export default function DeeplinkGenerator() {
  • For custom extensions:
    • Provide a unique ID, name, and description
    • -
    • Enter the command used to run your extension
    • +
    • Choose STDIO and enter the command to run your extension (e.g. npx @gooseai/my-ext), or choose Streamable HTTP and enter the endpoint URL
    • Add any required environment variables
    • +
    • For Streamable HTTP, add any required request headers
  • Click "Generate Deeplink" to create your installation deeplink.
  • -
  • Copy and share the generated deeplink - when users click it, it will open Goose Desktop and prompt them to install your extension.
  • +
  • Copy and share the generated deeplink — when users click it, it will open goose Desktop and prompt them to install your extension.
  • ); -} \ No newline at end of file +} From ed1dedc8320a60344a64ce8da5e8461c9d12d2f2 Mon Sep 17 00:00:00 2001 From: Jude Edwards Date: Tue, 23 Jun 2026 13:19:14 -0700 Subject: [PATCH 07/12] fix(cost): resolve databricks_v2 pricing and surface cost in standard usage update (#9925) Co-authored-by: Douwe M Osinga --- .../src/canonical/name_builder.rs | 53 ++++++++++++++++++- crates/goose/src/acp/server.rs | 10 +++- 2 files changed, 60 insertions(+), 3 deletions(-) diff --git a/crates/goose-providers/src/canonical/name_builder.rs b/crates/goose-providers/src/canonical/name_builder.rs index 6408201017dd..2d377bcf4272 100644 --- a/crates/goose-providers/src/canonical/name_builder.rs +++ b/crates/goose-providers/src/canonical/name_builder.rs @@ -35,7 +35,10 @@ pub fn canonical_name(provider: &str, model: &str) -> String { } fn is_meta_provider(provider: &str) -> bool { - matches!(provider, "databricks" | "tetrate" | "bedrock" | "azure") + matches!( + provider, + "databricks" | "databricks_v2" | "tetrate" | "bedrock" | "azure" + ) } pub fn map_provider_name(provider: &str) -> &str { @@ -46,6 +49,7 @@ pub fn map_provider_name(provider: &str) -> &str { "aws_bedrock" => "amazon-bedrock", "gcp_vertex_ai" => "google-vertex", "gemini_oauth" => "google", + "databricks_v2" => "databricks", "zhipu" => "zhipuai", "novita" => "novita-ai", "opencode_go" => "opencode-go", @@ -121,6 +125,21 @@ pub fn map_to_canonical_model( } } + // Fallback for meta-providers: some native aliases are keyed under the + // meta-provider itself (e.g. "databricks/databricks-gpt-oss-120b") and do + // not infer back to a first-party provider. Only try this after inference, + // so models that DO infer (e.g. databricks-claude-* -> anthropic/*) keep + // resolving to the richer first-party catalog entry. + if is_meta_provider(provider) { + if let Some(canonical) = registry.get(registry_provider, model) { + return Some(canonical.id.clone()); + } + let normalized_model = strip_version_suffix(model); + if let Some(canonical) = registry.get(registry_provider, &normalized_model) { + return Some(canonical.id.clone()); + } + } + None } @@ -537,4 +556,36 @@ mod tests { Some("google-vertex/claude-haiku-4.5".to_string()) ); } + + // Databricks-native open-weight ids are keyed under the meta-provider itself + // (e.g. "databricks/databricks-gpt-oss-120b") and do not infer back to another + // provider, so they must resolve via the direct meta-provider lookup. These + // particular ids are unversioned, so the assertions are not catalog-version brittle. + #[test] + fn test_databricks_native_open_weight_ids_resolve() { + let r = super::super::CanonicalModelRegistry::bundled().unwrap(); + + assert_eq!( + map_to_canonical_model("databricks_v2", "databricks-gpt-oss-120b", r), + Some("databricks/databricks-gpt-oss-120b".to_string()) + ); + assert_eq!( + map_to_canonical_model("databricks_v2", "databricks-gpt-oss-20b", r), + Some("databricks/databricks-gpt-oss-20b".to_string()) + ); + // Legacy provider name resolves identically. + assert_eq!( + map_to_canonical_model("databricks", "databricks-gpt-oss-120b", r), + Some("databricks/databricks-gpt-oss-120b".to_string()) + ); + + // Regression guard: the meta-provider lookup must remain a *fallback* + // after inference. databricks-claude-* aliases infer back to the richer + // first-party "anthropic/*" entry (which carries thinking_mode used for + // adaptive thinking), not the metadata-poor "databricks/databricks-*" one. + assert_eq!( + map_to_canonical_model("databricks", "databricks-claude-opus-4-7", r), + Some("anthropic/claude-opus-4.7".to_string()) + ); + } } diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 4b9b0c6af868..5b038cf03511 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -42,7 +42,7 @@ use crate::utils::sanitize_unicode_tags; use agent_client_protocol::schema::{ AgentCapabilities, Annotations, AuthMethod, AuthMethodAgent, AuthenticateRequest, AuthenticateResponse, BlobResourceContents, CancelNotification, CloseSessionRequest, - CloseSessionResponse, ConfigOptionUpdate, Content, ContentBlock, ContentChunk, + CloseSessionResponse, ConfigOptionUpdate, Content, ContentBlock, ContentChunk, Cost, CurrentModeUpdate, EmbeddedResource, EmbeddedResourceResource, FileSystemCapabilities, ForkSessionRequest, ForkSessionResponse, ImageContent, Implementation, InitializeRequest, InitializeResponse, ListSessionsRequest, ListSessionsResponse, LoadSessionRequest, @@ -834,7 +834,13 @@ pub(super) fn build_usage_updates(session: &Session) -> Option { accumulated_cost: session.accumulated_cost, }), }, - standard: UsageUpdate::new(used, ctx_limit), + standard: { + let mut standard = UsageUpdate::new(used, ctx_limit); + if let Some(amount) = session.accumulated_cost { + standard = standard.cost(Cost::new(amount, "USD")); + } + standard + }, }) } From f5a11718f2408c292e723df7ed6e08626721aa94 Mon Sep 17 00:00:00 2001 From: Jack Amadeo Date: Tue, 23 Jun 2026 19:18:07 -0400 Subject: [PATCH 08/12] provider refactor: don't require a model config to create a provider (#9953) Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- crates/goose-cli/src/commands/configure.rs | 45 ++- crates/goose-cli/src/commands/info.rs | 3 +- .../src/scenario_tests/scenario_runner.rs | 10 +- .../goose-cli/src/scenario_tests/scenarios.rs | 2 +- crates/goose-cli/src/session/builder.rs | 113 ++++---- crates/goose-cli/src/session/mod.rs | 52 ++-- crates/goose-providers/examples/streaming.rs | 4 +- crates/goose-providers/src/base.rs | 76 ++--- crates/goose-providers/src/formats/openai.rs | 5 +- .../src/formats/openai_responses.rs | 20 +- crates/goose-providers/src/model.rs | 271 +++++------------- crates/goose-providers/src/openai.rs | 80 +++--- .../goose-providers/src/openai_compatible.rs | 17 +- crates/goose-server/src/routes/agent.rs | 16 +- .../src/routes/config_management.rs | 31 +- crates/goose-server/src/routes/errors.rs | 7 - crates/goose-server/src/routes/sampling.rs | 8 +- crates/goose/examples/agent.rs | 9 +- crates/goose/examples/databricks_oauth.rs | 6 +- crates/goose/examples/image_tool.rs | 29 +- crates/goose/src/acp/provider.rs | 248 ++++++++++++---- crates/goose/src/acp/response_builder.rs | 33 +-- crates/goose/src/acp/server.rs | 65 ++++- crates/goose/src/acp/server/dispatch.rs | 4 +- crates/goose/src/acp/server/providers.rs | 52 ++-- crates/goose/src/acp/server_factory.rs | 12 +- crates/goose/src/agents/agent.rs | 259 ++++++++++------- crates/goose/src/agents/execute_commands.rs | 9 +- crates/goose/src/agents/mcp_client.rs | 15 +- crates/goose/src/agents/moim.rs | 33 ++- .../src/agents/platform_extensions/apps.rs | 4 +- .../src/agents/platform_extensions/mod.rs | 21 ++ .../platform_extensions/orchestrator.rs | 54 ++-- .../agents/platform_extensions/summarize.rs | 14 +- .../src/agents/platform_extensions/summon.rs | 30 +- crates/goose/src/agents/reply_parts.rs | 47 ++- crates/goose/src/agents/subagent_handler.rs | 6 +- .../goose/src/agents/subagent_task_config.rs | 3 + .../goose/src/bin/build_canonical_models.rs | 6 +- .../goose/src/config/declarative_providers.rs | 16 +- crates/goose/src/context_mgmt/mod.rs | 83 ++++-- crates/goose/src/doctor.rs | 61 ++-- crates/goose/src/execution/manager.rs | 19 +- crates/goose/src/model_config.rs | 90 +++++- .../src/permission/permission_inspector.rs | 28 +- .../goose/src/permission/permission_judge.rs | 31 +- crates/goose/src/providers/amp_acp.rs | 9 +- crates/goose/src/providers/anthropic.rs | 61 +--- crates/goose/src/providers/avian.rs | 3 - crates/goose/src/providers/azure.rs | 3 - crates/goose/src/providers/base.rs | 5 +- crates/goose/src/providers/bedrock.rs | 108 ++++--- crates/goose/src/providers/chatgpt_codex.rs | 24 +- crates/goose/src/providers/claude_acp.rs | 9 +- crates/goose/src/providers/claude_code.rs | 36 +-- crates/goose/src/providers/codex.rs | 48 +--- crates/goose/src/providers/codex_acp.rs | 9 +- crates/goose/src/providers/copilot_acp.rs | 23 +- crates/goose/src/providers/cursor_agent.rs | 22 +- crates/goose/src/providers/databricks.rs | 32 +-- crates/goose/src/providers/databricks_v2.rs | 13 +- .../goose/src/providers/formats/anthropic.rs | 11 +- crates/goose/src/providers/formats/bedrock.rs | 12 +- .../goose/src/providers/formats/databricks.rs | 19 +- crates/goose/src/providers/formats/google.rs | 18 +- .../goose/src/providers/formats/openrouter.rs | 10 +- .../goose/src/providers/formats/snowflake.rs | 6 +- crates/goose/src/providers/gcpvertexai.rs | 25 +- crates/goose/src/providers/gemini_cli.rs | 11 +- crates/goose/src/providers/gemini_oauth.rs | 17 +- crates/goose/src/providers/githubcopilot.rs | 10 +- crates/goose/src/providers/google.rs | 17 +- crates/goose/src/providers/huggingface.rs | 27 +- crates/goose/src/providers/init.rs | 44 +-- crates/goose/src/providers/inventory/mod.rs | 8 +- crates/goose/src/providers/kimicode.rs | 17 +- crates/goose/src/providers/litellm.rs | 48 ++-- crates/goose/src/providers/local_inference.rs | 13 +- crates/goose/src/providers/nanogpt.rs | 10 +- crates/goose/src/providers/ollama.rs | 40 +-- crates/goose/src/providers/openai_def.rs | 147 +++++----- crates/goose/src/providers/openrouter.rs | 31 +- crates/goose/src/providers/pi_acp.rs | 9 +- .../goose/src/providers/provider_registry.rs | 63 ++-- crates/goose/src/providers/provider_test.rs | 5 +- crates/goose/src/providers/sagemaker_tgi.rs | 31 +- crates/goose/src/providers/snowflake.rs | 12 +- crates/goose/src/providers/testprovider.rs | 17 +- crates/goose/src/providers/tetrate.rs | 10 +- crates/goose/src/providers/toolshim.rs | 10 +- crates/goose/src/providers/xai.rs | 3 - crates/goose/src/providers/xai_oauth.rs | 6 - crates/goose/src/scheduler.rs | 6 +- .../goose/src/security/adversary_inspector.rs | 57 +++- crates/goose/src/session/session_manager.rs | 35 ++- crates/goose/src/session/session_naming.rs | 13 +- crates/goose/tests/acp_common_tests/mod.rs | 2 +- .../goose/tests/acp_custom_requests_test.rs | 13 +- crates/goose/tests/acp_fixtures/mod.rs | 29 +- crates/goose/tests/acp_fixtures/provider.rs | 7 +- .../acp_secret_cache_invalidation_test.rs | 13 +- .../goose/tests/adversary_inspector_tests.rs | 24 +- crates/goose/tests/agent.rs | 80 +++--- crates/goose/tests/compaction.rs | 18 +- .../tests/local_inference_integration.rs | 18 +- crates/goose/tests/local_inference_perf.rs | 6 +- crates/goose/tests/mcp_integration_test.rs | 29 +- crates/goose/tests/providers.rs | 32 ++- .../tests/session_id_propagation_test.rs | 5 +- crates/goose/tests/tetrate_streaming.rs | 27 +- ui/desktop/openapi.json | 5 + ui/desktop/src/api/types.gen.ts | 5 + 112 files changed, 1819 insertions(+), 1764 deletions(-) diff --git a/crates/goose-cli/src/commands/configure.rs b/crates/goose-cli/src/commands/configure.rs index cf1dd5f04b83..d94b1d6e5e23 100644 --- a/crates/goose-cli/src/commands/configure.rs +++ b/crates/goose-cli/src/commands/configure.rs @@ -339,8 +339,7 @@ async fn handle_oauth_configuration(provider_name: &str, key_name: &str) -> anyh )); // Create a temporary provider instance to handle OAuth - let temp_model = goose::model_config::model_config_from_user_config(provider_name, "temp")?; - match create(provider_name, temp_model, Vec::new()).await { + match create(provider_name, Vec::new()).await { Ok(provider) => match provider.configure_oauth().await { Ok(_) => { let _ = cliclack::log::success("OAuth authentication completed successfully!"); @@ -761,13 +760,11 @@ pub async fn configure_provider_dialog() -> anyhow::Result { let spin = spinner(); spin.start("Attempting to fetch supported models..."); - let temp_model_config = goose::model_config::model_config_from_user_config( - provider_name, - &provider_meta.default_model, - )?; - let temp_provider = create(provider_name, temp_model_config, Vec::new()).await?; + let temp_provider = create(provider_name, Vec::new()).await?; let models_res = retry_operation(&RetryConfig::default(), || async { - temp_provider.fetch_recommended_models().await + temp_provider + .fetch_recommended_models(goose::model_config::global_toolshim()) + .await }) .await; spin.stop(style("Model fetch complete").green()); @@ -792,9 +789,7 @@ pub async fn configure_provider_dialog() -> anyhow::Result { { let supports_thinking = match temp_provider.fetch_model_info(&model).await { Ok(model_info) => model_info.reasoning, - Err(_) => goose_providers::model::ModelConfig::new(&model) - .map(|c| c.is_reasoning_model()) - .unwrap_or(false), + Err(_) => goose_providers::model::ModelConfig::new(&model).is_reasoning_model(), }; if supports_thinking { @@ -1618,8 +1613,10 @@ pub async fn configure_tool_permissions_dialog() -> anyhow::Result<()> { } let extensions = extension_config.into_iter().collect::>(); - let new_provider = create(&provider_name, model_config, extensions).await?; - agent.update_provider(new_provider, &session.id).await?; + let new_provider = create(&provider_name, extensions).await?; + agent + .update_provider(new_provider, model_config, &session.id) + .await?; let permission_manager = PermissionManager::instance(); let selected_tools = agent @@ -1811,12 +1808,11 @@ pub async fn handle_openrouter_auth() -> anyhow::Result<()> { } }; - match create("openrouter", model_config, Vec::new()).await { + match create("openrouter", Vec::new()).await { Ok(provider) => { - let provider_model_config = provider.get_model_config(); let test_result = provider .complete( - &provider_model_config, + &model_config, "", "You are goose, an AI assistant.", &[Message::user().with_text("Say 'Configuration test successful!'")], @@ -1883,17 +1879,14 @@ pub async fn handle_tetrate_auth() -> anyhow::Result<()> { // Test configuration println!("\nTesting configuration..."); let configured_model: String = config.get_goose_model()?; - let model_config = - match goose::model_config::model_config_from_user_config("tetrate", &configured_model) { - Ok(config) => config, - Err(e) => { - eprintln!("⚠️ Invalid model configuration: {}", e); - eprintln!("Your settings have been saved. Please check your model configuration."); - return Ok(()); - } - }; + if let Err(e) = goose::model_config::model_config_from_user_config("tetrate", &configured_model) + { + eprintln!("⚠️ Invalid model configuration: {}", e); + eprintln!("Your settings have been saved. Please check your model configuration."); + return Ok(()); + } - match create("tetrate", model_config, Vec::new()).await { + match create("tetrate", Vec::new()).await { Ok(provider) => { let test_result = provider.fetch_supported_models().await; diff --git a/crates/goose-cli/src/commands/info.rs b/crates/goose-cli/src/commands/info.rs index 447d1a8ed88a..d21771d29cb5 100644 --- a/crates/goose-cli/src/commands/info.rs +++ b/crates/goose-cli/src/commands/info.rs @@ -76,7 +76,7 @@ async fn check_provider( let model_config = goose::model_config::model_config_from_user_config(&provider, &model) .map_err(|e| ProviderCheckError::InvalidModel(e.to_string()))?; - let provider_client = goose::providers::create(&provider, model_config, Vec::new()) + let provider_client = goose::providers::create(&provider, Vec::new()) .await .map_err(|e| { let error = e.to_string(); @@ -87,7 +87,6 @@ async fn check_provider( })?; let test_msg = Message::user().with_text("Say 'ok'"); - let model_config = provider_client.get_model_config(); let start = std::time::Instant::now(); provider_client .complete(&model_config, "check", "", &[test_msg], &[]) diff --git a/crates/goose-cli/src/scenario_tests/scenario_runner.rs b/crates/goose-cli/src/scenario_tests/scenario_runner.rs index 9619f3bdc720..2498c78cc0d3 100644 --- a/crates/goose-cli/src/scenario_tests/scenario_runner.rs +++ b/crates/goose-cli/src/scenario_tests/scenario_runner.rs @@ -185,12 +185,7 @@ where let original_env = setup_environment(config)?; - let inner_provider = create( - &factory_name, - goose::model_config::model_config_from_user_config(&factory_name, config.model_name)?, - Vec::new(), - ) - .await?; + let inner_provider = create(&factory_name, Vec::new()).await?; let test_provider = Arc::new(TestProvider::new_recording(inner_provider, &file_path)); ( @@ -245,9 +240,12 @@ where ) .await?; + let scenario_model_config = + goose::model_config::model_config_from_user_config(&factory_name, config.model_name)?; agent .update_provider( provider_arc as Arc, + scenario_model_config, &session.id, ) .await?; diff --git a/crates/goose-cli/src/scenario_tests/scenarios.rs b/crates/goose-cli/src/scenario_tests/scenarios.rs index d1c7695742cc..c52df9753fce 100644 --- a/crates/goose-cli/src/scenario_tests/scenarios.rs +++ b/crates/goose-cli/src/scenario_tests/scenarios.rs @@ -80,7 +80,7 @@ mod tests { // run_scenario( // "context_length_exceeded", // Box::new(|provider| { - // let model_config = provider.get_model_config(); + // let model_config = resolve_global_model_config(); // let context_length = model_config.context_limit.unwrap_or(300_000); // // "hello " is only one token in most models, since the hello and space often // // occur together in the training data. diff --git a/crates/goose-cli/src/session/builder.rs b/crates/goose-cli/src/session/builder.rs index 5639a01fce99..2823d374f8a0 100644 --- a/crates/goose-cli/src/session/builder.rs +++ b/crates/goose-cli/src/session/builder.rs @@ -544,84 +544,79 @@ pub async fn build_session(session_config: SessionBuilderConfig) -> CliSession { } }; - let (new_provider, effective_provider_name, effective_model_name) = match create( - &resolved.provider_name, - resolved.model_config.clone(), - extensions_for_provider.clone(), - ) - .await - { - Ok(provider) => ( - provider, - resolved.provider_name.clone(), - resolved.model_name.clone(), - ), - Err(e) - if session_config.resume - && session_config.provider.is_none() - && is_provider_unavailable_error(&e) => - { - let fallback_provider = config.get_goose_provider().unwrap_or_else(|_| { - output::render_error("No provider configured. Run 'goose configure' first."); - process::exit(1); - }); - let fallback_model = config.get_goose_model().unwrap_or_else(|_| { - output::render_error("No model configured. Run 'goose configure' first."); - process::exit(1); - }); - eprintln!( - "{}", - style(format!( - "Warning: Could not create the session's original provider '{}' ({}). \ + let (new_provider, effective_provider_name, effective_model_name, effective_model_config) = + match create(&resolved.provider_name, extensions_for_provider.clone()).await { + Ok(provider) => ( + provider, + resolved.provider_name.clone(), + resolved.model_name.clone(), + resolved.model_config.clone(), + ), + Err(e) + if session_config.resume + && session_config.provider.is_none() + && is_provider_unavailable_error(&e) => + { + let fallback_provider = config.get_goose_provider().unwrap_or_else(|_| { + output::render_error("No provider configured. Run 'goose configure' first."); + process::exit(1); + }); + let fallback_model = config.get_goose_model().unwrap_or_else(|_| { + output::render_error("No model configured. Run 'goose configure' first."); + process::exit(1); + }); + eprintln!( + "{}", + style(format!( + "Warning: Could not create the session's original provider '{}' ({}). \ Falling back to the default provider '{}'.", - resolved.provider_name, e, fallback_provider - )) - .yellow() - ); - let fallback_model_config = - model_config_from_user_config(fallback_provider.as_str(), &fallback_model) - .unwrap_or_else(|e| { + resolved.provider_name, e, fallback_provider + )) + .yellow() + ); + let fallback_model_config = + model_config_from_user_config(fallback_provider.as_str(), &fallback_model) + .unwrap_or_else(|e| { + output::render_error(&format!( + "Failed to create model configuration: {}", + e + )); + process::exit(1); + }); + match create(&fallback_provider, extensions_for_provider.clone()).await { + Ok(provider) => ( + provider, + fallback_provider, + fallback_model, + fallback_model_config, + ), + Err(e2) => { output::render_error(&format!( - "Failed to create model configuration: {}", - e - )); - process::exit(1); - }); - match create( - &fallback_provider, - fallback_model_config, - extensions_for_provider.clone(), - ) - .await - { - Ok(provider) => (provider, fallback_provider, fallback_model), - Err(e2) => { - output::render_error(&format!( "Error {}.\n\ Please check your system keychain and run 'goose configure' again.\n\ If your system is unable to use the keyring, please try setting secret key(s) via environment variables.\n\ For more info, see: https://goose-docs.ai/docs/troubleshooting/#keychainkeyring-errors", e2 )); - process::exit(1); + process::exit(1); + } } } - } - Err(e) => { - output::render_error(&format!( + Err(e) => { + output::render_error(&format!( "Error {}.\n\ Please check your system keychain and run 'goose configure' again.\n\ If your system is unable to use the keyring, please try setting secret key(s) via environment variables.\n\ For more info, see: https://goose-docs.ai/docs/troubleshooting/#keychainkeyring-errors", e )); - process::exit(1); - } - }; + process::exit(1); + } + }; tracing::info!("🤖 Using model: {}", effective_model_name); agent - .update_provider(new_provider, &session_id) + .update_provider(new_provider, effective_model_config, &session_id) .await .unwrap_or_else(|e| { output::render_error(&format!("Failed to initialize agent: {}", e)); diff --git a/crates/goose-cli/src/session/mod.rs b/crates/goose-cli/src/session/mod.rs index 76d7ac91513f..49c9d50ec497 100644 --- a/crates/goose-cli/src/session/mod.rs +++ b/crates/goose-cli/src/session/mod.rs @@ -219,13 +219,13 @@ pub async fn classify_planner_response( session_id: &str, message_text: String, provider: Arc, + model_config: goose_providers::model::ModelConfig, ) -> Result { let prompt = format!( "The text below is the output from an AI model which can either provide a plan or list of clarifying questions. Based on the text below, decide if the output is a \"plan\" or \"clarifying questions\".\n---\n{message_text}" ); let message = Message::user().with_text(&prompt); - let model_config = provider.get_model_config(); let (result, _usage) = provider .complete( &model_config, @@ -722,8 +722,8 @@ impl CliSession { RunMode::Plan => { let mut plan_messages = self.messages.clone(); plan_messages.push(Message::user().with_text(content)); - let reasoner = get_reasoner().await?; - self.plan_with_reasoner_model(plan_messages, reasoner) + let (reasoner, reasoner_model_config) = get_reasoner().await?; + self.plan_with_reasoner_model(plan_messages, reasoner, reasoner_model_config) .await?; } } @@ -810,7 +810,10 @@ impl CliSession { async fn handle_model(&self, model: Option<&str>) -> Result<()> { let provider = self.agent.provider().await?; let current_provider_name = provider.get_name().to_string(); - let current_model_config = provider.get_model_config(); + let current_model_config = self + .agent + .model_config_for_session(&self.session_id) + .await?; let current_model_name = current_model_config.model_name.clone(); if model.is_none() { @@ -859,13 +862,12 @@ impl CliSession { } let extensions = self.agent.get_extension_configs().await; - let new_provider = - goose::providers::create(¤t_provider_name, new_model_config, extensions) - .await - .map_err(|e| anyhow::anyhow!("Failed to create provider: {e}"))?; + let new_provider = goose::providers::create(¤t_provider_name, extensions) + .await + .map_err(|e| anyhow::anyhow!("Failed to create provider: {e}"))?; self.agent - .update_provider(new_provider, &self.session_id) + .update_provider(new_provider, new_model_config, &self.session_id) .await?; let mode = self.agent.goose_mode().await; @@ -888,8 +890,9 @@ impl CliSession { let mut plan_messages = self.messages.clone(); plan_messages.push(Message::user().with_text(&options.message_text)); - let reasoner = get_reasoner().await?; - self.plan_with_reasoner_model(plan_messages, reasoner).await + let (reasoner, reasoner_model_config) = get_reasoner().await?; + self.plan_with_reasoner_model(plan_messages, reasoner, reasoner_model_config) + .await } async fn handle_clear(&mut self) -> Result<()> { @@ -1051,10 +1054,10 @@ impl CliSession { &mut self, plan_messages: Conversation, reasoner: Arc, + model_config: goose_providers::model::ModelConfig, ) -> Result<(), anyhow::Error> { let plan_prompt = self.agent.get_plan_prompt(&self.session_id).await?; output::show_thinking(); - let model_config = reasoner.get_model_config(); let (plan_response, _usage) = reasoner .complete( &model_config, @@ -1070,6 +1073,9 @@ impl CliSession { &self.session_id, plan_response.as_concat_text(), self.agent.provider().await?, + self.agent + .model_config_for_session(&self.session_id) + .await?, ) .await?; @@ -1596,8 +1602,14 @@ impl CliSession { /// Display enhanced context usage with session totals pub async fn display_context_usage(&self) -> Result<()> { let provider = self.agent.provider().await?; - let model_config = provider.get_model_config(); - let context_limit = model_config.context_limit(); + let model_config = self + .agent + .model_config_for_session(&self.session_id) + .await?; + let context_limit = provider + .get_context_limit(&model_config) + .await + .unwrap_or_else(|_| model_config.context_limit()); let config = Config::global(); let show_cost = config @@ -2213,7 +2225,8 @@ fn handle_agent_error(e: &anyhow::Error, is_stream_json_mode: bool) { } } -async fn get_reasoner() -> Result, anyhow::Error> { +async fn get_reasoner( +) -> Result<(Arc, goose_providers::model::ModelConfig), anyhow::Error> { use goose::providers::create; let config = Config::global(); @@ -2252,9 +2265,9 @@ async fn get_reasoner() -> Result, anyhow::Error> { goose::model_config::model_config_from_user_config(&provider, model.as_str())? .with_context_limit(planner_context_limit); let extensions = goose::config::extensions::get_enabled_extensions_with_config(config); - let reasoner = create(&provider, model_config, extensions).await?; + let reasoner = create(&provider, extensions).await?; - Ok(reasoner) + Ok((reasoner, model_config)) } /// Format elapsed time duration @@ -2434,7 +2447,6 @@ mod tests { max_tokens: Some(16_384), toolshim: true, toolshim_model: Some("qwen2.5-coder".to_string()), - fast_model_config: None, request_params: Some(HashMap::from([( "anthropic_beta".to_string(), serde_json::json!(["output-128k-2025-02-19"]), @@ -2444,7 +2456,7 @@ mod tests { let switched = build_switched_model_config("openai", "gpt-5.4", ¤t_model_config).unwrap(); - let expected = goose_providers::model::ModelConfig::new_or_fail("gpt-5.4") + let expected = goose_providers::model::ModelConfig::new("gpt-5.4") .with_canonical_limits("openai") .with_temperature(Some(0.25)) .with_toolshim(true) @@ -2471,7 +2483,7 @@ mod tests { ("GOOSE_THINKING_EFFORT", None::<&str>), ]); - let current = goose_providers::model::ModelConfig::new_or_fail("gpt-5.4-high") + let current = goose_providers::model::ModelConfig::new("gpt-5.4-high") .with_canonical_limits("openai"); assert_eq!(current.model_name, "gpt-5.4"); assert_eq!( diff --git a/crates/goose-providers/examples/streaming.rs b/crates/goose-providers/examples/streaming.rs index 1775cb1d69c9..efb482be3a5f 100644 --- a/crates/goose-providers/examples/streaming.rs +++ b/crates/goose-providers/examples/streaming.rs @@ -12,18 +12,18 @@ use goose_providers::{ #[tokio::main] async fn main() -> Result<()> { - let model = ModelConfig::new("gpt-5.4-mini")?; let key = env::var("OPENAI_API_KEY").map_err(|_| anyhow::anyhow!("need an OpenAI key"))?; let api_client = ApiClient::new_with_tls( "https://api.openai.com".to_string(), AuthMethod::BearerToken(key), Some(Default::default()), )?; - let provider = OpenAiProvider::new(api_client, model.clone()); + let provider = OpenAiProvider::new(api_client); let system = "You are a knowledgable geography expert"; let messages = [Message::user().with_text("what is the capital of France?")]; + let model = ModelConfig::new("gpt-5.4-mini"); let mut stream = provider .stream( &model, diff --git a/crates/goose-providers/src/base.rs b/crates/goose-providers/src/base.rs index 8ee3034303b7..855aebd7908f 100644 --- a/crates/goose-providers/src/base.rs +++ b/crates/goose-providers/src/base.rs @@ -41,6 +41,10 @@ pub struct ProviderMetadata { /// 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, + /// The name of a fast/cheap model to use for lightweight tasks (e.g. session naming, + /// compaction). When set, fast-path callers prefer this model over the main model. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub fast_model: Option, } impl ProviderMetadata { @@ -66,6 +70,7 @@ impl ProviderMetadata { config_keys, setup_steps: vec![], model_selection_hint: None, + fast_model: None, } } @@ -88,6 +93,7 @@ impl ProviderMetadata { config_keys, setup_steps: vec![], model_selection_hint: None, + fast_model: None, } } @@ -102,6 +108,7 @@ impl ProviderMetadata { config_keys: vec![], setup_steps: vec![], model_selection_hint: None, + fast_model: None, } } @@ -114,6 +121,11 @@ impl ProviderMetadata { self.model_selection_hint = Some(hint.to_string()); self } + + pub fn with_fast_model(mut self, fast_model: &str) -> Self { + self.fast_model = Some(fast_model.to_string()); + self + } } /// Configuration key metadata for provider setup @@ -291,12 +303,12 @@ pub fn model_info_for_provider_model(provider_name: &str, model_name: &str) -> M let reasoning = canonical .as_ref() .and_then(|model| model.reasoning) - .unwrap_or_else(|| ModelConfig::new_or_fail(model_name).is_reasoning_model()); + .unwrap_or_else(|| ModelConfig::new(model_name).is_reasoning_model()); ModelInfo { name: model_name.to_string(), resolved_model: None, - context_limit: ModelConfig::new_or_fail(model_name) + context_limit: ModelConfig::new(model_name) .with_canonical_limits(provider_name) .context_limit(), input_token_cost: None, @@ -403,43 +415,15 @@ pub trait Provider: Send + Sync { collect_stream(stream).await } - /// Try fast model first, fall back to regular model on failure. - async fn complete_fast( - &self, - session_id: &str, - system: &str, - messages: &[Message], - tools: &[Tool], - ) -> Result<(Message, ProviderUsage), ProviderError> { - let model_config = self.get_model_config(); - let fast_config = model_config.use_fast_model(); - - let result = self - .complete(&fast_config, session_id, system, messages, tools) - .await; - - match result { - Ok(response) => Ok(response), - Err(e) => { - if fast_config.model_name != model_config.model_name { - tracing::warn!( - "Fast model {} failed with error: {}. Falling back to regular model {}", - fast_config.model_name, - e, - model_config.model_name - ); - self.complete(&model_config, session_id, system, messages, tools) - .await - } else { - Err(e) - } - } - } + /// Resolve the effective context limit for a model config. + /// + /// Providers may override this to enrich the limit with provider-specific + /// metadata (e.g. cached model info or a value captured from a remote + /// session). The default returns the limit derived from the model config. + async fn get_context_limit(&self, model_config: &ModelConfig) -> Result { + Ok(model_config.context_limit()) } - /// Get the model config from the provider - fn get_model_config(&self) -> ModelConfig; - fn retry_config(&self) -> RetryConfig { RetryConfig::default() } @@ -466,7 +450,10 @@ pub trait Provider: Send + Sync { } /// Fetch inventory models filtered by canonical registry and usability. - async fn fetch_recommended_models(&self) -> Result, ProviderError> { + /// + /// When `toolshim` is true, models that lack native tool-call support are + /// retained because the toolshim layer emulates tool calling. + async fn fetch_recommended_models(&self, toolshim: bool) -> Result, ProviderError> { let all_models = self.fetch_supported_models().await?; if self.skip_canonical_filtering() { @@ -496,7 +483,7 @@ pub trait Provider: Send + Sync { return None; } - if !canonical_model.tool_call && !self.get_model_config().toolshim { + if !canonical_model.tool_call && !toolshim { return None; } @@ -526,9 +513,12 @@ pub trait Provider: Send + Sync { } } - async fn fetch_recommended_model_info(&self) -> Result, ProviderError> { + async fn fetch_recommended_model_info( + &self, + toolshim: bool, + ) -> Result, ProviderError> { Ok(self - .fetch_recommended_models() + .fetch_recommended_models(toolshim) .await? .iter() .map(|model_name| model_info_for_provider_model(self.get_name(), model_name)) @@ -558,10 +548,6 @@ pub trait Provider: Send + Sync { false } - async fn supports_cache_control(&self) -> bool { - false - } - /// Configure OAuth authentication for this provider /// /// This method is called when a provider has configuration keys marked with oauth_flow = true. diff --git a/crates/goose-providers/src/formats/openai.rs b/crates/goose-providers/src/formats/openai.rs index a824e56a6653..43d96992f45a 100644 --- a/crates/goose-providers/src/formats/openai.rs +++ b/crates/goose-providers/src/formats/openai.rs @@ -1452,10 +1452,7 @@ mod tests { use tokio_stream::{self, StreamExt}; fn test_model_config(model_name: &str) -> ModelConfig { - ModelConfig { - model_name: model_name.to_string(), - ..Default::default() - } + ModelConfig::new(model_name) } #[test] diff --git a/crates/goose-providers/src/formats/openai_responses.rs b/crates/goose-providers/src/formats/openai_responses.rs index 38a3023c5cc1..54cbc0b09cbb 100644 --- a/crates/goose-providers/src/formats/openai_responses.rs +++ b/crates/goose-providers/src/formats/openai_responses.rs @@ -1200,7 +1200,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1292,7 +1291,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1336,7 +1334,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1358,7 +1355,7 @@ mod tests { #[test] fn test_responses_request_with_normalized_effort_suffix() { - let model_config = ModelConfig::new("o3-mini-high").unwrap(); + let model_config = ModelConfig::new("o3-mini-high"); let result = create_responses_request(&model_config, "You are helpful.", &[], &[]).unwrap(); @@ -1377,7 +1374,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1403,7 +1399,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1432,7 +1427,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1479,7 +1473,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1516,7 +1509,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1548,7 +1540,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1579,7 +1570,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1614,7 +1604,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1651,7 +1640,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1677,7 +1665,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1709,7 +1696,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1741,7 +1727,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1863,7 +1848,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1901,7 +1885,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1939,7 +1922,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; diff --git a/crates/goose-providers/src/model.rs b/crates/goose-providers/src/model.rs index 22159af5df08..cb5ed56e52d3 100644 --- a/crates/goose-providers/src/model.rs +++ b/crates/goose-providers/src/model.rs @@ -4,22 +4,11 @@ use serde::de::Deserializer; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::collections::HashMap; -use thiserror::Error; use utoipa::ToSchema; pub const DEFAULT_CONTEXT_LIMIT: usize = 128_000; -#[derive(Error, Debug)] -pub enum ConfigError { - #[error("Environment variable '{0}' not found")] - EnvVarMissing(String), - #[error("Invalid value for '{0}': '{1}' - {2}")] - InvalidValue(String, String, String), - #[error("Value for '{0}' is out of valid range: {1}")] - InvalidRange(String, String), -} - -#[derive(Debug, Clone, Default, Serialize, ToSchema)] +#[derive(Debug, Clone, Serialize, ToSchema)] pub struct ModelConfig { pub model_name: String, pub context_limit: Option, @@ -27,8 +16,6 @@ pub struct ModelConfig { pub max_tokens: Option, pub toolshim: bool, pub toolshim_model: Option, - #[serde(skip)] - pub fast_model_config: Option>, /// Provider-specific request parameters (e.g., anthropic_beta headers) #[serde(default, skip_serializing_if = "Option::is_none")] pub request_params: Option>, @@ -49,8 +36,6 @@ impl<'de> Deserialize<'de> for ModelConfig { max_tokens: Option, toolshim: bool, toolshim_model: Option, - #[serde(default)] - fast_model_config: Option>, #[serde(default, skip_serializing_if = "Option::is_none")] request_params: Option>, #[serde(default, skip_serializing_if = "Option::is_none")] @@ -65,7 +50,6 @@ impl<'de> Deserialize<'de> for ModelConfig { max_tokens: raw.max_tokens, toolshim: raw.toolshim, toolshim_model: raw.toolshim_model, - fast_model_config: raw.fast_model_config, request_params: raw.request_params, reasoning: raw.reasoning, }; @@ -75,7 +59,7 @@ impl<'de> Deserialize<'de> for ModelConfig { } impl ModelConfig { - pub fn new(model_name: impl AsRef) -> Result { + pub fn new(model_name: impl AsRef) -> Self { let mut config = Self { model_name: model_name.as_ref().to_string(), context_limit: None, @@ -83,12 +67,11 @@ impl ModelConfig { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; config.normalize_effort_suffix(); - Ok(config) + config } pub fn with_canonical_limits(mut self, provider_name: &str) -> Self { @@ -147,11 +130,6 @@ impl ModelConfig { if self.context_limit.is_none() { self.context_limit = limit; } - - if let Some(fast_config) = self.fast_model_config.take() { - self.fast_model_config = Some(Box::new(fast_config.with_default_context_limit(limit))); - } - self } @@ -159,11 +137,6 @@ impl ModelConfig { if self.max_tokens.is_none() { self.max_tokens = tokens; } - - if let Some(fast_config) = self.fast_model_config.take() { - self.fast_model_config = Some(Box::new(fast_config.with_default_max_tokens(tokens))); - } - self } @@ -177,21 +150,6 @@ impl ModelConfig { self } - pub fn with_fast( - mut self, - fast_model_name: &str, - provider_name: &str, - ) -> Result { - let fast_config = ModelConfig::new(fast_model_name)?.with_canonical_limits(provider_name); - self.fast_model_config = Some(Box::new(fast_config)); - Ok(self) - } - - pub fn with_fast_model_config(mut self, fast_model_config: ModelConfig) -> Self { - self.fast_model_config = Some(Box::new(fast_model_config)); - self - } - pub fn with_merged_request_params(mut self, params: HashMap) -> Self { match self.request_params.as_mut() { Some(existing) => { @@ -221,12 +179,6 @@ impl ModelConfig { self = self.with_thinking_effort(effort); } } - - if let Some(fast_config) = self.fast_model_config.take() { - self.fast_model_config = - Some(Box::new(fast_config.with_default_thinking_effort(effort))); - } - self } @@ -262,14 +214,6 @@ impl ModelConfig { self } - pub fn use_fast_model(&self) -> Self { - if let Some(fast_config) = &self.fast_model_config { - *fast_config.clone() - } else { - self.clone() - } - } - pub fn context_limit(&self) -> usize { self.context_limit.unwrap_or(DEFAULT_CONTEXT_LIMIT) } @@ -347,66 +291,30 @@ impl ModelConfig { .and_then(|params| params.get(request_key)) .and_then(|v| serde_json::from_value(v.clone()).ok()) } - - pub fn new_or_fail(model_name: &str) -> ModelConfig { - ModelConfig::new(model_name) - .unwrap_or_else(|_| panic!("Failed to create model config for {}", model_name)) - } } #[cfg(test)] mod tests { use super::*; - #[test] - fn test_deserialize_preserves_fast_model_config() { - let config: ModelConfig = serde_json::from_value(serde_json::json!({ - "model_name": "primary-model", - "context_limit": null, - "temperature": null, - "max_tokens": null, - "toolshim": false, - "toolshim_model": null, - "fast_model_config": { - "model_name": "fast-model", - "context_limit": 4096, - "temperature": null, - "max_tokens": 1024, - "toolshim": false, - "toolshim_model": null - } - })) - .unwrap(); - - let fast_config = config.fast_model_config.as_ref().unwrap(); - assert_eq!(fast_config.model_name, "fast-model"); - assert_eq!(fast_config.context_limit, Some(4096)); - assert_eq!(fast_config.max_tokens, Some(1024)); - assert_eq!(config.use_fast_model().model_name, "fast-model"); - } - mod thinking_effort_tests { use super::*; + fn config_with_params(model_name: &str, params: HashMap) -> ModelConfig { + ModelConfig::new(model_name).with_merged_request_params(params) + } + #[test] fn from_request_params() { let mut params = HashMap::new(); params.insert("thinking_effort".to_string(), serde_json::json!("medium")); - let config = ModelConfig { - model_name: "test".to_string(), - request_params: Some(params), - ..Default::default() - }; + let config = config_with_params("test", params); assert_eq!(config.thinking_effort(), Some(ThinkingEffort::Medium)); } #[test] fn with_thinking_effort_sets_request_param() { - let config = ModelConfig { - model_name: "test".to_string(), - ..Default::default() - } - .with_thinking_effort(ThinkingEffort::High); + let config = ModelConfig::new("test").with_thinking_effort(ThinkingEffort::High); assert_eq!( config @@ -419,19 +327,12 @@ mod tests { #[test] fn preserves_explicit_thinking_effort() { - let previous = ModelConfig { - model_name: "previous".to_string(), - request_params: Some(HashMap::from([( - "thinking_effort".to_string(), - serde_json::json!("high"), - )])), - ..Default::default() - }; - let config = ModelConfig { - model_name: "next".to_string(), - ..Default::default() - } - .with_inherited_session_settings_from(Some(&previous), None); + let previous = config_with_params( + "previous", + HashMap::from([("thinking_effort".to_string(), serde_json::json!("high"))]), + ); + let config = ModelConfig::new("next") + .with_inherited_session_settings_from(Some(&previous), None); assert_eq!( config @@ -444,22 +345,14 @@ mod tests { #[test] fn does_not_override_existing_thinking_effort() { - let previous = ModelConfig { - model_name: "previous".to_string(), - request_params: Some(HashMap::from([( - "thinking_effort".to_string(), - serde_json::json!("high"), - )])), - ..Default::default() - }; - let config = ModelConfig { - model_name: "next".to_string(), - request_params: Some(HashMap::from([( - "thinking_effort".to_string(), - serde_json::json!("low"), - )])), - ..Default::default() - } + let previous = config_with_params( + "previous", + HashMap::from([("thinking_effort".to_string(), serde_json::json!("high"))]), + ); + let config = config_with_params( + "next", + HashMap::from([("thinking_effort".to_string(), serde_json::json!("low"))]), + ) .with_inherited_session_settings_from(Some(&previous), None); assert_eq!( @@ -473,38 +366,23 @@ mod tests { #[test] fn does_not_preserve_unrelated_request_params() { - let previous = ModelConfig { - model_name: "previous".to_string(), - request_params: Some(HashMap::from([( - "provider_specific".to_string(), - serde_json::json!("old"), - )])), - ..Default::default() - }; - let config = ModelConfig { - model_name: "next".to_string(), - ..Default::default() - } - .with_inherited_session_settings_from(Some(&previous), None); + let previous = config_with_params( + "previous", + HashMap::from([("provider_specific".to_string(), serde_json::json!("old"))]), + ); + let config = ModelConfig::new("next") + .with_inherited_session_settings_from(Some(&previous), None); assert!(config.request_params.is_none()); } #[test] fn explicit_request_params_override_preserved_session_settings() { - let previous = ModelConfig { - model_name: "previous".to_string(), - request_params: Some(HashMap::from([( - "thinking_effort".to_string(), - serde_json::json!("high"), - )])), - ..Default::default() - }; - let config = ModelConfig { - model_name: "next".to_string(), - ..Default::default() - } - .with_inherited_session_settings_from( + let previous = config_with_params( + "previous", + HashMap::from([("thinking_effort".to_string(), serde_json::json!("high"))]), + ); + let config = ModelConfig::new("next").with_inherited_session_settings_from( Some(&previous), Some(HashMap::from([( "thinking_effort".to_string(), @@ -531,7 +409,7 @@ mod tests { ("GOOSE_TOOLSHIM", None::<&str>), ("GOOSE_TOOLSHIM_OLLAMA_MODEL", None::<&str>), ]); - let config = ModelConfig::new("o3-mini-high").unwrap(); + let config = ModelConfig::new("o3-mini-high"); assert_eq!(config.model_name, "o3-mini"); assert_eq!(config.thinking_effort(), Some(ThinkingEffort::High)); } @@ -546,7 +424,7 @@ mod tests { ("GOOSE_TOOLSHIM", None::<&str>), ("GOOSE_TOOLSHIM_OLLAMA_MODEL", None::<&str>), ]); - let config = ModelConfig::new("o3-mini-none").unwrap(); + let config = ModelConfig::new("o3-mini-none"); assert_eq!(config.model_name, "o3-mini"); assert_eq!(config.thinking_effort(), Some(ThinkingEffort::Off)); } @@ -561,7 +439,7 @@ mod tests { ("GOOSE_TOOLSHIM", None::<&str>), ("GOOSE_TOOLSHIM_OLLAMA_MODEL", None::<&str>), ]); - let config = ModelConfig::new("gpt-5.4-xhigh").unwrap(); + let config = ModelConfig::new("gpt-5.4-xhigh"); assert_eq!(config.model_name, "gpt-5.4"); assert_eq!(config.thinking_effort(), Some(ThinkingEffort::Max)); } @@ -578,7 +456,7 @@ mod tests { ]); let mut params = HashMap::new(); params.insert("thinking_effort".to_string(), serde_json::json!("low")); - let mut config = ModelConfig::new("o3-mini-high").unwrap(); + let mut config = ModelConfig::new("o3-mini-high"); // Suffix was already normalized during new(), but if request_params // were set before construction, the suffix would not be stripped. // Verify the normalized state: @@ -599,7 +477,7 @@ mod tests { ("GOOSE_TOOLSHIM", None::<&str>), ("GOOSE_TOOLSHIM_OLLAMA_MODEL", None::<&str>), ]); - let config = ModelConfig::new("o3-mini").unwrap(); + let config = ModelConfig::new("o3-mini"); assert_eq!(config.model_name, "o3-mini"); } @@ -613,7 +491,7 @@ mod tests { ("GOOSE_TOOLSHIM", None::<&str>), ("GOOSE_TOOLSHIM_OLLAMA_MODEL", None::<&str>), ]); - let config = ModelConfig::new("claude-sonnet-4-high").unwrap(); + let config = ModelConfig::new("claude-sonnet-4-high"); assert_eq!(config.model_name, "claude-sonnet-4-high"); } @@ -640,7 +518,7 @@ mod tests { ("GOOSE_MAX_TOKENS", None::<&str>), ("GOOSE_CONTEXT_LIMIT", None::<&str>), ]); - let config = ModelConfig::new_or_fail("gpt-4o").with_canonical_limits("openai"); + let config = ModelConfig::new("gpt-4o").with_canonical_limits("openai"); assert_eq!(config.context_limit, Some(128_000)); assert_eq!(config.max_tokens, Some(16_384)); @@ -653,7 +531,7 @@ mod tests { ("GOOSE_MAX_TOKENS", None::<&str>), ("GOOSE_CONTEXT_LIMIT", None::<&str>), ]); - let mut config = ModelConfig::new_or_fail("gpt-4o"); + let mut config = ModelConfig::new("gpt-4o"); config.context_limit = Some(64_000); let config = config.with_canonical_limits("openai"); @@ -666,7 +544,7 @@ mod tests { ("GOOSE_MAX_TOKENS", None::<&str>), ("GOOSE_CONTEXT_LIMIT", None::<&str>), ]); - let mut config = ModelConfig::new_or_fail("gpt-4o"); + let mut config = ModelConfig::new("gpt-4o"); config.max_tokens = Some(1_000); let config = config.with_canonical_limits("openai"); @@ -679,8 +557,7 @@ mod tests { ("GOOSE_MAX_TOKENS", None::<&str>), ("GOOSE_CONTEXT_LIMIT", None::<&str>), ]); - let config = - ModelConfig::new_or_fail("moonshotai/kimi-k2.6").with_canonical_limits("nvidia"); + let config = ModelConfig::new("moonshotai/kimi-k2.6").with_canonical_limits("nvidia"); assert_eq!(config.context_limit, Some(262_144)); assert_eq!(config.max_tokens, None); @@ -693,8 +570,7 @@ mod tests { ("GOOSE_MAX_TOKENS", None::<&str>), ("GOOSE_CONTEXT_LIMIT", None::<&str>), ]); - let config = - ModelConfig::new_or_fail("totally-unknown-model").with_canonical_limits("openai"); + let config = ModelConfig::new("totally-unknown-model").with_canonical_limits("openai"); assert_eq!(config.context_limit, None); assert_eq!(config.max_tokens, None); @@ -709,17 +585,16 @@ mod tests { ]); // "databricks-gpt-5.4-high" should resolve via "databricks-gpt-5.4" - let config = ModelConfig::new_or_fail("databricks-gpt-5.4-high") - .with_canonical_limits("databricks"); + let config = + ModelConfig::new("databricks-gpt-5.4-high").with_canonical_limits("databricks"); assert_eq!(config.context_limit, Some(1_050_000)); // "gpt-5.4-xhigh" should resolve via "gpt-5.4" - let config = ModelConfig::new_or_fail("gpt-5.4-xhigh").with_canonical_limits("openai"); + let config = ModelConfig::new("gpt-5.4-xhigh").with_canonical_limits("openai"); assert_eq!(config.context_limit, Some(1_050_000)); // "gpt-5.4-nano-low" should resolve via "gpt-5.4-nano" - let config = - ModelConfig::new_or_fail("gpt-5.4-nano-low").with_canonical_limits("openai"); + let config = ModelConfig::new("gpt-5.4-nano-low").with_canonical_limits("openai"); assert_eq!(config.context_limit, Some(400_000)); } } @@ -738,41 +613,39 @@ mod tests { #[test] fn bare_reasoning_models() { let _guard = env_lock::lock_env(ENV_LOCK_KEYS); - assert!(ModelConfig::new_or_fail("o1").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("o1-preview").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("o3").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("o3-mini").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("o4-mini").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("gpt-5").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("gpt-5-3-codex").is_openai_reasoning_model()); + assert!(ModelConfig::new("o1").is_openai_reasoning_model()); + assert!(ModelConfig::new("o1-preview").is_openai_reasoning_model()); + assert!(ModelConfig::new("o3").is_openai_reasoning_model()); + assert!(ModelConfig::new("o3-mini").is_openai_reasoning_model()); + assert!(ModelConfig::new("o4-mini").is_openai_reasoning_model()); + assert!(ModelConfig::new("gpt-5").is_openai_reasoning_model()); + assert!(ModelConfig::new("gpt-5-3-codex").is_openai_reasoning_model()); } #[test] fn goose_prefixed_reasoning_models() { let _guard = env_lock::lock_env(ENV_LOCK_KEYS); - assert!(ModelConfig::new_or_fail("goose-o3-mini").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("goose-o4-mini").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("goose-gpt-5").is_openai_reasoning_model()); + assert!(ModelConfig::new("goose-o3-mini").is_openai_reasoning_model()); + assert!(ModelConfig::new("goose-o4-mini").is_openai_reasoning_model()); + assert!(ModelConfig::new("goose-gpt-5").is_openai_reasoning_model()); } #[test] fn databricks_prefixed_reasoning_models() { let _guard = env_lock::lock_env(ENV_LOCK_KEYS); - assert!(ModelConfig::new_or_fail("databricks-o3-mini").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("databricks-o4-mini").is_openai_reasoning_model()); - assert!(ModelConfig::new_or_fail("databricks-gpt-5").is_openai_reasoning_model()); + assert!(ModelConfig::new("databricks-o3-mini").is_openai_reasoning_model()); + assert!(ModelConfig::new("databricks-o4-mini").is_openai_reasoning_model()); + assert!(ModelConfig::new("databricks-gpt-5").is_openai_reasoning_model()); } #[test] fn non_reasoning_models() { let _guard = env_lock::lock_env(ENV_LOCK_KEYS); - assert!(!ModelConfig::new_or_fail("claude-sonnet-4").is_openai_reasoning_model()); - assert!(!ModelConfig::new_or_fail("gpt-4o").is_openai_reasoning_model()); - assert!( - !ModelConfig::new_or_fail("databricks-claude-sonnet-4").is_openai_reasoning_model() - ); - assert!(!ModelConfig::new_or_fail("goose-claude-sonnet-4").is_openai_reasoning_model()); - assert!(!ModelConfig::new_or_fail("llama-3-70b").is_openai_reasoning_model()); + assert!(!ModelConfig::new("claude-sonnet-4").is_openai_reasoning_model()); + assert!(!ModelConfig::new("gpt-4o").is_openai_reasoning_model()); + assert!(!ModelConfig::new("databricks-claude-sonnet-4").is_openai_reasoning_model()); + assert!(!ModelConfig::new("goose-claude-sonnet-4").is_openai_reasoning_model()); + assert!(!ModelConfig::new("llama-3-70b").is_openai_reasoning_model()); } } @@ -790,19 +663,19 @@ mod tests { #[test] fn includes_reasoning_model_families() { let _guard = env_lock::lock_env(ENV_LOCK_KEYS); - assert!(ModelConfig::new_or_fail("o3-mini").is_reasoning_model()); - assert!(ModelConfig::new_or_fail("claude-sonnet-4").is_reasoning_model()); - assert!(ModelConfig::new_or_fail("gemini-3-pro").is_reasoning_model()); + assert!(ModelConfig::new("o3-mini").is_reasoning_model()); + assert!(ModelConfig::new("claude-sonnet-4").is_reasoning_model()); + assert!(ModelConfig::new("gemini-3-pro").is_reasoning_model()); } #[test] fn uses_explicit_metadata_first() { let _guard = env_lock::lock_env(ENV_LOCK_KEYS); - let mut config = ModelConfig::new_or_fail("provider-alias"); + let mut config = ModelConfig::new("provider-alias"); config.reasoning = Some(true); assert!(config.is_reasoning_model()); - let mut config = ModelConfig::new_or_fail("claude-sonnet-4"); + let mut config = ModelConfig::new("claude-sonnet-4"); config.reasoning = Some(false); assert!(!config.is_reasoning_model()); } diff --git a/crates/goose-providers/src/openai.rs b/crates/goose-providers/src/openai.rs index f0a915eb7439..9a9719cc6fc5 100644 --- a/crates/goose-providers/src/openai.rs +++ b/crates/goose-providers/src/openai.rs @@ -20,6 +20,7 @@ use anyhow::Result; use async_trait::async_trait; use reqwest::StatusCode; use std::collections::HashMap; +use std::sync::{Arc, Mutex}; use crate::base::{MessageStream, ProviderDescriptor}; use crate::model::ModelConfig; @@ -125,7 +126,6 @@ pub struct OpenAiProvider { base_path: String, organization: Option, project: Option, - model: ModelConfig, custom_headers: Option>, supports_streaming: bool, name: String, @@ -133,6 +133,8 @@ pub struct OpenAiProvider { dynamic_models: Option, skip_canonical_filtering: bool, preserve_thinking_context: bool, + #[serde(skip)] + n_ctx_cache: Arc>>>, } /// Builder for [`OpenAiProvider`]. @@ -145,7 +147,6 @@ pub struct OpenAiProviderBuilder { base_path: String, organization: Option, project: Option, - model: ModelConfig, custom_headers: Option>, supports_streaming: bool, name: String, @@ -156,13 +157,12 @@ pub struct OpenAiProviderBuilder { } impl OpenAiProviderBuilder { - pub fn new(api_client: ApiClient, model: ModelConfig) -> Self { + pub fn new(api_client: ApiClient) -> Self { Self { api_client, base_path: OPEN_AI_DEFAULT_BASE_PATH.to_string(), organization: None, project: None, - model, custom_headers: None, supports_streaming: true, name: OPEN_AI_PROVIDER_NAME.to_string(), @@ -193,11 +193,6 @@ impl OpenAiProviderBuilder { self } - pub fn model(mut self, model: ModelConfig) -> Self { - self.model = model; - self - } - pub fn custom_headers(mut self, custom_headers: Option>) -> Self { self.custom_headers = custom_headers; self @@ -239,7 +234,6 @@ impl OpenAiProviderBuilder { base_path: self.base_path, organization: self.organization, project: self.project, - model: self.model, custom_headers: self.custom_headers, supports_streaming: self.supports_streaming, name: self.name, @@ -247,19 +241,19 @@ impl OpenAiProviderBuilder { dynamic_models: self.dynamic_models, skip_canonical_filtering: self.skip_canonical_filtering, preserve_thinking_context: self.preserve_thinking_context, + n_ctx_cache: Arc::new(Mutex::new(HashMap::new())), } } } impl OpenAiProvider { #[doc(hidden)] - pub fn new(api_client: ApiClient, model: ModelConfig) -> Self { + pub fn new(api_client: ApiClient) -> Self { Self { api_client, base_path: OPEN_AI_DEFAULT_BASE_PATH.to_string(), organization: None, project: None, - model, custom_headers: None, supports_streaming: true, name: OPEN_AI_PROVIDER_NAME.to_string(), @@ -267,6 +261,7 @@ impl OpenAiProvider { dynamic_models: None, skip_canonical_filtering: false, preserve_thinking_context: false, + n_ctx_cache: Arc::new(Mutex::new(HashMap::new())), } } @@ -390,28 +385,6 @@ impl OpenAiProvider { } } - /// Fill the model's context limit from the API when it isn't already set. - /// - /// An existing value may be an explicit GOOSE_CONTEXT_LIMIT, an ACP/server - /// per-session override, or a GOOSE_PREDEFINED_MODELS entry, none of which we - /// should overwrite. llama.cpp and Ollama report the real allocated window via - /// the non-standard meta.n_ctx field; reading it fixes auto-compaction for local - /// servers that would otherwise fall back to DEFAULT_CONTEXT_LIMIT. The probe is - /// bounded by a short timeout so a hung /v1/models can't stall provider - /// construction (the shared ApiClient uses OPENAI_TIMEOUT, up to 600s). - pub async fn probe_context_limit_if_unset(&mut self) { - if self.model.context_limit.is_some() { - return; - } - const N_CTX_PROBE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5); - let model_name = self.model.model_name.clone(); - if let Ok(Some(n_ctx)) = - tokio::time::timeout(N_CTX_PROBE_TIMEOUT, self.fetch_n_ctx_from_api(&model_name)).await - { - self.model.context_limit = Some(n_ctx); - } - } - async fn fetch_models_from_api(&self) -> Result, ProviderError> { let models_path = Self::map_base_path(&self.base_path, "models", OPEN_AI_DEFAULT_MODELS_PATH); @@ -546,8 +519,41 @@ impl Provider for OpenAiProvider { self.skip_canonical_filtering } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() + /// Resolve the effective context limit. When the config carries an explicit + /// limit (GOOSE_CONTEXT_LIMIT, a session override, or a known/canonical + /// value) it is used as-is. Otherwise probe `/v1/models`: llama.cpp and + /// Ollama report the real allocated window via the non-standard + /// `meta.n_ctx` field, which fixes auto-compaction for local servers that + /// would otherwise fall back to DEFAULT_CONTEXT_LIMIT. The probe is bounded + /// by a short timeout so a hung endpoint can't stall the caller. + async fn get_context_limit(&self, model_config: &ModelConfig) -> Result { + if let Some(limit) = model_config.context_limit { + return Ok(limit); + } + + if let Some(cached) = self + .n_ctx_cache + .lock() + .ok() + .and_then(|cache| cache.get(&model_config.model_name).copied()) + { + return Ok(cached.unwrap_or_else(|| model_config.context_limit())); + } + + const N_CTX_PROBE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5); + let probed = tokio::time::timeout( + N_CTX_PROBE_TIMEOUT, + self.fetch_n_ctx_from_api(&model_config.model_name), + ) + .await + .ok() + .flatten(); + + if let Ok(mut cache) = self.n_ctx_cache.lock() { + cache.insert(model_config.model_name.clone(), probed); + } + + Ok(probed.unwrap_or_else(|| model_config.context_limit())) } async fn fetch_supported_models(&self) -> Result, ProviderError> { @@ -715,7 +721,6 @@ mod tests { base_path: "v1/chat/completions".to_string(), organization: None, project: None, - model: ModelConfig::new_or_fail("test-model"), custom_headers: None, supports_streaming: true, name: name.to_string(), @@ -723,6 +728,7 @@ mod tests { dynamic_models: None, skip_canonical_filtering: false, preserve_thinking_context: false, + n_ctx_cache: Arc::new(Mutex::new(HashMap::new())), } } diff --git a/crates/goose-providers/src/openai_compatible.rs b/crates/goose-providers/src/openai_compatible.rs index cc872f13c98c..68aad564a606 100644 --- a/crates/goose-providers/src/openai_compatible.rs +++ b/crates/goose-providers/src/openai_compatible.rs @@ -29,23 +29,16 @@ pub struct OpenAiCompatibleProvider { name: String, /// Client targeted at the base URL (e.g. `https://api.x.ai/v1`) api_client: ApiClient, - model: ModelConfig, /// Path prefix prepended to `chat/completions` (e.g. `"deployments/{name}/"` for Azure). completions_prefix: String, supports_streaming: bool, } impl OpenAiCompatibleProvider { - pub fn new( - name: String, - api_client: ApiClient, - model: ModelConfig, - completions_prefix: String, - ) -> Self { + pub fn new(name: String, api_client: ApiClient, completions_prefix: String) -> Self { Self { name, api_client, - model, completions_prefix, supports_streaming: true, } @@ -82,10 +75,6 @@ impl Provider for OpenAiCompatibleProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client @@ -312,13 +301,13 @@ mod tests { None, ) .unwrap(), - ModelConfig::new_or_fail("test-model"), String::new(), ) .with_supports_streaming(false); + let model = ModelConfig::new("test-model"); let payload = provider - .build_request(&provider.model, "", &[], &[], provider.supports_streaming) + .build_request(&model, "", &[], &[], provider.supports_streaming) .unwrap(); assert_eq!(payload.get("stream"), None); diff --git a/crates/goose-server/src/routes/agent.rs b/crates/goose-server/src/routes/agent.rs index 06cf11652575..c93f6b1596f5 100644 --- a/crates/goose-server/src/routes/agent.rs +++ b/crates/goose-server/src/routes/agent.rs @@ -660,17 +660,15 @@ async fn update_agent_provider( EnabledExtensionsState::for_session(state.session_manager(), &payload.session_id, config) .await; - let new_provider = create(&payload.provider, model_config, extensions) - .await - .map_err(|e| { - ( - StatusCode::BAD_REQUEST, - format!("Failed to create {} provider: {}", &payload.provider, e), - ) - })?; + let new_provider = create(&payload.provider, extensions).await.map_err(|e| { + ( + StatusCode::BAD_REQUEST, + format!("Failed to create {} provider: {}", &payload.provider, e), + ) + })?; agent - .update_provider(new_provider, &payload.session_id) + .update_provider(new_provider, model_config, &payload.session_id) .await .map_err(|e| { ( diff --git a/crates/goose-server/src/routes/config_management.rs b/crates/goose-server/src/routes/config_management.rs index 6cae9dad1bd4..4ad7f189fa25 100644 --- a/crates/goose-server/src/routes/config_management.rs +++ b/crates/goose-server/src/routes/config_management.rs @@ -937,11 +937,11 @@ pub async fn get_provider_models( ))); } - let model_config = - goose::model_config::model_config_from_user_config(&name, &metadata.default_model)?; - let provider = goose::providers::create(&name, model_config, Vec::new()).await?; + let provider = goose::providers::create(&name, Vec::new()).await?; - let models_result = provider.fetch_recommended_model_info().await; + let models_result = provider + .fetch_recommended_model_info(goose::model_config::global_toolshim()) + .await; match models_result { Ok(models) => Ok(Json(models)), @@ -973,7 +973,7 @@ pub async fn resolve_provider_model_info( } let model_config = goose::model_config::model_config_from_user_config(name, model)?; - let provider = goose::providers::create(name, model_config.clone(), Vec::new()).await?; + let provider = goose::providers::create(name, Vec::new()).await?; match provider.fetch_model_info(model).await { Ok(info) => Ok(info), Err(error) => { @@ -1124,7 +1124,7 @@ pub async fn get_canonical_model_info( max_output_tokens: canonical_model.limit.output, reasoning: canonical_model .reasoning - .unwrap_or_else(|| ModelConfig::new_or_fail(&query.model).is_reasoning_model()), + .unwrap_or_else(|| ModelConfig::new(&query.model).is_reasoning_model()), // Costs are per million tokens - client handles division for display input_token_cost: canonical_model.cost.input, output_token_cost: canonical_model.cost.output, @@ -1440,20 +1440,13 @@ pub async fn configure_provider_oauth( return Ok(Json("OAuth configuration completed".to_string())); } - let temp_model = goose::model_config::model_config_from_user_config(&provider_name, "temp") - .map_err(|e| { - ErrorResponse::bad_request(format!("Failed to create temporary model config: {}", e)) - })?; - // OAuth configuration does not use extensions. - let provider = create(&provider_name, temp_model, Vec::new()) - .await - .map_err(|e| { - ErrorResponse::bad_request(format!( - "Failed to create provider '{}': {}", - provider_name, e - )) - })?; + let provider = create(&provider_name, Vec::new()).await.map_err(|e| { + ErrorResponse::bad_request(format!( + "Failed to create provider '{}': {}", + provider_name, e + )) + })?; provider.configure_oauth().await.map_err(|e| { ErrorResponse::bad_request(format!( diff --git a/crates/goose-server/src/routes/errors.rs b/crates/goose-server/src/routes/errors.rs index 9d3de5a373dd..ab78df291ad3 100644 --- a/crates/goose-server/src/routes/errors.rs +++ b/crates/goose-server/src/routes/errors.rs @@ -5,7 +5,6 @@ use axum::{ }; use goose::config::ConfigError; use goose_providers::errors::ProviderError; -use goose_providers::model::ConfigError as ModelConfigError; use serde::Serialize; use utoipa::ToSchema; @@ -77,12 +76,6 @@ impl From for ErrorResponse { } } -impl From for ErrorResponse { - fn from(err: ModelConfigError) -> Self { - Self::internal(format!("Model configuration error: {}", err)) - } -} - impl From for ErrorResponse { fn from(status: StatusCode) -> Self { let message = status.canonical_reason().unwrap_or("Unknown error"); diff --git a/crates/goose-server/src/routes/sampling.rs b/crates/goose-server/src/routes/sampling.rs index fb92c9f0ad25..63c74b956a3b 100644 --- a/crates/goose-server/src/routes/sampling.rs +++ b/crates/goose-server/src/routes/sampling.rs @@ -51,7 +51,13 @@ async fn create_message( .as_deref() .unwrap_or("You are a helpful AI assistant."); - let model_config = provider.get_model_config(); + let model_config = agent + .model_config_for_session(&session_id) + .await + .map_err(|e| { + tracing::error!("Failed to resolve model config: {}", e); + StatusCode::INTERNAL_SERVER_ERROR + })?; let (response, usage) = provider .complete(&model_config, &session_id, system, &messages, &[]) .await diff --git a/crates/goose/examples/agent.rs b/crates/goose/examples/agent.rs index c28f1df0f74b..8a045304bc38 100644 --- a/crates/goose/examples/agent.rs +++ b/crates/goose/examples/agent.rs @@ -12,8 +12,9 @@ use std::path::PathBuf; async fn main() -> anyhow::Result<()> { let _ = dotenv(); - let provider = - create_with_named_model("databricks", DATABRICKS_DEFAULT_MODEL, Vec::new()).await?; + let provider = create_with_named_model("databricks", Vec::new()).await?; + let model_config = + goose::model_config::model_config_from_user_config("databricks", DATABRICKS_DEFAULT_MODEL)?; let agent = Agent::new(); @@ -28,7 +29,9 @@ async fn main() -> anyhow::Result<()> { ) .await?; - agent.update_provider(provider, &session.id).await?; + agent + .update_provider(provider, model_config, &session.id) + .await?; let config = ExtensionConfig::stdio( "developer", diff --git a/crates/goose/examples/databricks_oauth.rs b/crates/goose/examples/databricks_oauth.rs index ce7b53842c1a..7e8a34bb748c 100644 --- a/crates/goose/examples/databricks_oauth.rs +++ b/crates/goose/examples/databricks_oauth.rs @@ -10,12 +10,12 @@ async fn main() -> Result<()> { std::env::remove_var("DATABRICKS_TOKEN"); - let provider = - create_with_named_model("databricks", DATABRICKS_DEFAULT_MODEL, Vec::new()).await?; + let provider = create_with_named_model("databricks", Vec::new()).await?; let message = Message::user().with_text("Tell me a short joke about programming."); - let model_config = provider.get_model_config(); + let model_config = + goose::model_config::model_config_from_user_config("databricks", DATABRICKS_DEFAULT_MODEL)?; let (response, usage) = provider .complete( &model_config, diff --git a/crates/goose/examples/image_tool.rs b/crates/goose/examples/image_tool.rs index f679a217bd61..b6f2949665c5 100644 --- a/crates/goose/examples/image_tool.rs +++ b/crates/goose/examples/image_tool.rs @@ -17,12 +17,30 @@ async fn main() -> Result<()> { dotenv().ok(); // Create providers - let providers: Vec> = vec![ - create_with_named_model("databricks", DATABRICKS_DEFAULT_MODEL, Vec::new()).await?, - create_with_named_model("openai", OPEN_AI_DEFAULT_MODEL, Vec::new()).await?, - create_with_named_model("anthropic", ANTHROPIC_DEFAULT_MODEL, Vec::new()).await?, + let providers: Vec<( + Arc, + goose_providers::model::ModelConfig, + )> = vec![ + ( + create_with_named_model("databricks", Vec::new()).await?, + goose::model_config::model_config_from_user_config( + "databricks", + DATABRICKS_DEFAULT_MODEL, + )?, + ), + ( + create_with_named_model("openai", Vec::new()).await?, + goose::model_config::model_config_from_user_config("openai", OPEN_AI_DEFAULT_MODEL)?, + ), + ( + create_with_named_model("anthropic", Vec::new()).await?, + goose::model_config::model_config_from_user_config( + "anthropic", + ANTHROPIC_DEFAULT_MODEL, + )?, + ), ]; - for provider in providers { + for (provider, model_config) in providers { // Read and encode test image let image_data = fs::read("crates/goose/examples/test_assets/test_image.png")?; let base64_image = BASE64.encode(image_data); @@ -56,7 +74,6 @@ async fn main() -> Result<()> { }, } }); - let model_config = provider.get_model_config(); let (response, usage) = provider .complete( &model_config, diff --git a/crates/goose/src/acp/provider.rs b/crates/goose/src/acp/provider.rs index fe970a10f309..576d1a5e3e2e 100644 --- a/crates/goose/src/acp/provider.rs +++ b/crates/goose/src/acp/provider.rs @@ -52,6 +52,11 @@ pub struct AcpProviderConfig { pub work_dir: PathBuf, pub mcp_servers: Vec, pub session_mode_id: Option, + pub session_config_options: Vec<(String, String)>, + /// Config option id used to select the model (e.g. `"model"`). When set, the + /// provider re-applies this option from the per-completion `ModelConfig` + /// whenever the active session model changes. + pub model_config_option_id: Option, pub mode_mapping: HashMap, pub notification_callback: Option>, } @@ -135,7 +140,6 @@ struct HandoffContextClaim { pub struct AcpProvider { name: String, - model: ModelConfig, goose_mode: Arc>, mode_mapping: HashMap, @@ -147,10 +151,16 @@ pub struct AcpProvider { handoff_context_sent: AtomicBool, /// Latest `size` reported by the ACP server in a `session/update` → /// `usage_update` notification. 0 means no real update has arrived yet, - /// in which case `get_model_config()` falls back to the static model + /// in which case `get_context_limit()` falls back to the supplied model /// configuration's context limit. context_size: Arc, + /// Config option id used to select the model, if this agent supports it. + model_config_option_id: Option, + /// Model currently applied via `model_config_option_id`, used to avoid + /// redundant `SetConfigOption` calls. + applied_model: Arc>>, + tx: Option>, loop_thread: Option>, } @@ -159,7 +169,6 @@ impl std::fmt::Debug for AcpProvider { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("AcpProvider") .field("name", &self.name) - .field("model", &self.model) .finish() } } @@ -177,13 +186,11 @@ fn spawn_client_loop(fut: impl Future + Send + 'static) -> JoinHand impl AcpProvider { pub async fn connect( name: String, - model: ModelConfig, goose_mode: GooseMode, config: AcpProviderConfig, ) -> Result { Self::start( name, - model, goose_mode, config, Box::new(|cl, rx, init_tx| Box::pin(cl.spawn(rx, init_tx))), @@ -194,14 +201,12 @@ impl AcpProvider { #[doc(hidden)] pub async fn connect_with_transport( name: String, - model: ModelConfig, goose_mode: GooseMode, config: AcpProviderConfig, transport: impl agent_client_protocol::ConnectTo + 'static, ) -> Result { Self::start( name, - model, goose_mode, config, Box::new(move |cl, mut rx, init_tx| { @@ -217,7 +222,6 @@ impl AcpProvider { async fn start( name: String, - model: ModelConfig, goose_mode: GooseMode, config: AcpProviderConfig, run: ClientLoopFn, @@ -225,6 +229,14 @@ impl AcpProvider { let (tx, rx) = mpsc::channel(32); let (init_tx, init_rx) = oneshot::channel(); let mode_mapping = config.mode_mapping.clone(); + let model_config_option_id = config.model_config_option_id.clone(); + let applied_model = config.model_config_option_id.as_ref().and_then(|id| { + config + .session_config_options + .iter() + .find(|(opt_id, _)| opt_id == id) + .map(|(_, value)| value.clone()) + }); let goose_mode_shared = Arc::new(Mutex::new(goose_mode)); let pending_tool_updates: Arc>> = Arc::new(Mutex::new(HashMap::new())); @@ -252,21 +264,6 @@ impl AcpProvider { .await .context("ACP session creation cancelled")??; - // Resolve model from the session response. - let resolved_model = if model.model_name == ACP_CURRENT_MODEL { - if let Ok((resolved, _)) = resolve_model_info(&name, &response) { - tracing::info!(from = ACP_CURRENT_MODEL, to = %resolved, "resolved ACP model"); - ModelConfig { - model_name: resolved, - ..model - } - } else { - model - } - } else { - model - }; - let session = AcpSession { id: response.session_id.clone(), response, @@ -274,7 +271,6 @@ impl AcpProvider { Ok(Self { name, - model: resolved_model, goose_mode: goose_mode_shared, mode_mapping, session, @@ -282,6 +278,8 @@ impl AcpProvider { pending_tool_updates, handoff_context_sent: AtomicBool::new(false), context_size, + model_config_option_id, + applied_model: Arc::new(Mutex::new(applied_model)), tx: Some(tx), loop_thread: Some(loop_thread), }) @@ -329,6 +327,39 @@ impl AcpProvider { response_rx.await.context("ACP request cancelled")? } + /// Re-apply the model selection config option when the active session model + /// differs from what was last applied. ACP agents that select their model + /// via a config option (e.g. Copilot) need this so resumed or switched + /// sessions actually use the requested model instead of the agent default. + async fn apply_model_if_changed(&self, model_name: &str) -> Result<()> { + let Some(config_id) = self.model_config_option_id.clone() else { + return Ok(()); + }; + if model_name == ACP_CURRENT_MODEL { + return Ok(()); + } + + { + let applied = self + .applied_model + .lock() + .map_err(|_| anyhow::anyhow!("applied_model lock poisoned"))?; + if applied.as_deref() == Some(model_name) { + return Ok(()); + } + } + + self.send_set_config_option("", config_id, model_name.to_string()) + .await?; + + let mut applied = self + .applied_model + .lock() + .map_err(|_| anyhow::anyhow!("applied_model lock poisoned"))?; + *applied = Some(model_name.to_string()); + Ok(()) + } + async fn prompt( &self, session_id: SessionId, @@ -378,13 +409,12 @@ impl Provider for AcpProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - let mut model = self.model.clone(); + async fn get_context_limit(&self, model_config: &ModelConfig) -> Result { let size = self.context_size.load(Ordering::Relaxed); if size > 0 { - model.context_limit = Some(size as usize); + return Ok(size as usize); } - model + Ok(model_config.context_limit()) } async fn update_mode(&self, session_id: &str, mode: GooseMode) -> Result<(), ProviderError> { @@ -441,6 +471,12 @@ impl Provider for AcpProvider { ) -> Result { let session_id = self.acp_session_id(); + self.apply_model_if_changed(&model_config.model_name) + .await + .map_err(|e| { + ProviderError::RequestFailed(format!("Failed to set ACP model option: {e}")) + })?; + let claim = self.claim_handoff_context(messages); let prompt_blocks = messages_to_prompt(messages, claim.include_context); // Drop any tool-call buffer state left over from a prior prompt @@ -1049,6 +1085,8 @@ async fn handle_requests( let result = match session { Ok(session) => { session_ids.push(session.session_id.clone()); + apply_session_config_options(&config, &cx, session.session_id.clone()) + .await?; apply_session_mode(&config, &goose_mode, &cx, session).await } Err(err) => Err(anyhow::anyhow!( @@ -1140,6 +1178,31 @@ async fn handle_requests( Ok(()) } +async fn apply_session_config_options( + config: &AcpProviderConfig, + cx: &ConnectionTo, + session_id: SessionId, +) -> Result<()> { + for (config_id, value) in &config.session_config_options { + let value_id = agent_client_protocol::schema::SessionConfigValueId::new(value.clone()); + cx.send_request(SetSessionConfigOptionRequest::new( + session_id.clone(), + config_id.clone(), + value_id, + )) + .block_task() + .await + .map_err(|err| { + anyhow::anyhow!( + "ACP agent rejected {} for '{}': {err}", + AGENT_METHOD_NAMES.session_set_config_option, + config_id + ) + })?; + } + Ok(()) +} + async fn apply_session_mode( config: &AcpProviderConfig, goose_mode: &Arc>, @@ -1531,30 +1594,33 @@ mod tests { } } - fn test_provider() -> AcpProvider { + fn test_provider() -> (AcpProvider, ModelConfig) { test_provider_with_tx(None) } - fn test_provider_with_tx(tx: Option>) -> AcpProvider { - AcpProvider { - name: "acp-test".to_string(), - model: ModelConfig { - model_name: "test-model".to_string(), - ..Default::default() - }, - goose_mode: Arc::new(Mutex::new(GooseMode::Auto)), - mode_mapping: HashMap::new(), - session: AcpSession { - id: SessionId::new("test-session"), - response: NewSessionResponse::new("test-session"), + fn test_provider_with_tx( + tx: Option>, + ) -> (AcpProvider, ModelConfig) { + ( + AcpProvider { + name: "acp-test".to_string(), + goose_mode: Arc::new(Mutex::new(GooseMode::Auto)), + mode_mapping: HashMap::new(), + session: AcpSession { + id: SessionId::new("test-session"), + response: NewSessionResponse::new("test-session"), + }, + pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), + pending_tool_updates: Arc::new(Mutex::new(HashMap::new())), + handoff_context_sent: AtomicBool::new(false), + context_size: Arc::new(AtomicU64::new(0)), + model_config_option_id: None, + applied_model: Arc::new(Mutex::new(None)), + tx, + loop_thread: None, }, - pending_confirmations: Arc::new(TokioMutex::new(HashMap::new())), - pending_tool_updates: Arc::new(Mutex::new(HashMap::new())), - handoff_context_sent: AtomicBool::new(false), - context_size: Arc::new(AtomicU64::new(0)), - tx, - loop_thread: None, - } + ModelConfig::new("test-model"), + ) } #[test] @@ -1623,7 +1689,7 @@ mod tests { #[test] fn handoff_context_is_sent_only_on_first_provider_prompt() { - let provider = test_provider(); + let (provider, _) = test_provider(); let messages = vec![ Message::assistant().with_text("prior answer"), Message::user().with_text("current request"), @@ -1640,7 +1706,7 @@ mod tests { #[test] fn first_prompt_without_history_still_marks_handoff_context_sent() { - let provider = test_provider(); + let (provider, _) = test_provider(); let first_prompt = vec![Message::user().with_text("new conversation")]; let later_prompt_with_history = vec![ Message::assistant().with_text("prior answer"), @@ -1656,30 +1722,30 @@ mod tests { assert!(!later_claim.include_context); } - #[test] - fn get_model_config_surfaces_captured_context_size() { - let provider = test_provider(); + #[tokio::test] + async fn get_context_limit_surfaces_captured_context_size() { + let (provider, model) = test_provider(); assert_eq!( - provider.get_model_config().context_limit(), + provider.get_context_limit(&model).await.unwrap(), goose_providers::model::DEFAULT_CONTEXT_LIMIT ); provider.context_size.store(200_000, Ordering::Relaxed); - assert_eq!(provider.get_model_config().context_limit(), 200_000); + assert_eq!(provider.get_context_limit(&model).await.unwrap(), 200_000); } #[tokio::test] async fn failed_first_prompt_send_rolls_back_handoff_context_claim() { let (tx, rx) = mpsc::channel(1); drop(rx); - let provider = test_provider_with_tx(Some(tx)); + let (provider, model) = test_provider_with_tx(Some(tx)); let messages = vec![ Message::assistant().with_text("prior answer"), Message::user().with_text("current request"), ]; let result = provider - .stream(&provider.model, "goose-session", "", &messages, &[]) + .stream(&model, "goose-session", "", &messages, &[]) .await; assert!(matches!(result, Err(ProviderError::RequestFailed(_)))); @@ -1688,6 +1754,74 @@ mod tests { assert!(next_claim.include_context); } + fn test_provider_with_model_option( + tx: mpsc::Sender, + applied_model: Option, + ) -> AcpProvider { + let (mut provider, _) = test_provider_with_tx(Some(tx)); + provider.model_config_option_id = Some("model".to_string()); + provider.applied_model = Arc::new(Mutex::new(applied_model)); + provider + } + + #[tokio::test] + async fn apply_model_if_changed_sends_set_config_option_on_change() { + let (tx, mut rx) = mpsc::channel(1); + let provider = test_provider_with_model_option(tx, Some("old-model".to_string())); + + let handle = + tokio::spawn(async move { provider.apply_model_if_changed("new-model").await }); + + match rx.recv().await.expect("expected a SetConfigOption request") { + ClientRequest::SetConfigOption { + config_id, + value, + response_tx, + .. + } => { + assert_eq!(config_id, "model"); + assert_eq!(value, "new-model"); + let _ = response_tx.send(Ok(())); + } + _ => panic!("unexpected request kind"), + } + + handle.await.unwrap().unwrap(); + } + + #[tokio::test] + async fn apply_model_if_changed_skips_when_model_unchanged() { + let (tx, mut rx) = mpsc::channel(1); + let provider = test_provider_with_model_option(tx, Some("same-model".to_string())); + + provider.apply_model_if_changed("same-model").await.unwrap(); + + assert!(rx.try_recv().is_err()); + } + + #[tokio::test] + async fn apply_model_if_changed_noop_without_option_id() { + let (tx, mut rx) = mpsc::channel(1); + let (provider, _) = test_provider_with_tx(Some(tx)); + + provider.apply_model_if_changed("any-model").await.unwrap(); + + assert!(rx.try_recv().is_err()); + } + + #[tokio::test] + async fn apply_model_if_changed_skips_sentinel_model() { + let (tx, mut rx) = mpsc::channel(1); + let provider = test_provider_with_model_option(tx, None); + + provider + .apply_model_if_changed(ACP_CURRENT_MODEL) + .await + .unwrap(); + + assert!(rx.try_recv().is_err()); + } + #[test] fn messages_to_prompt_includes_all_prior_handoff_context() { let messages = vec![ diff --git a/crates/goose/src/acp/response_builder.rs b/crates/goose/src/acp/response_builder.rs index cb1115f4ac28..d307d911529c 100644 --- a/crates/goose/src/acp/response_builder.rs +++ b/crates/goose/src/acp/response_builder.rs @@ -550,14 +550,11 @@ mod tests { provider_options: Vec, model_state: SessionModelState, ) -> Vec { - let model_config = ModelConfig { - model_name: model_state.current_model_id.0.to_string(), - request_params: Some(std::collections::HashMap::from([( + let model_config = ModelConfig::new(model_state.current_model_id.0.as_ref()) + .with_merged_request_params(std::collections::HashMap::from([( "thinking_effort".to_string(), serde_json::json!("off"), - )])), - ..Default::default() - }; + )])); build_config_options( &mode_state, &model_state, @@ -577,14 +574,12 @@ mod tests { "claude-sonnet-4", )], ); - let model_config = ModelConfig { - model_name: "claude-sonnet-4".to_string(), - request_params: Some(std::collections::HashMap::from([( + let model_config = ModelConfig::new("claude-sonnet-4").with_merged_request_params( + std::collections::HashMap::from([( "thinking_effort".to_string(), serde_json::json!("high"), - )])), - ..Default::default() - }; + )]), + ); let options = build_config_options( &mode_state, @@ -612,15 +607,11 @@ mod tests { ModelId::new("gpt-4"), vec![ModelInfo::new(ModelId::new("gpt-4"), "gpt-4")], ); - let model_config = ModelConfig { - model_name: "gpt-4".to_string(), - request_params: Some(std::collections::HashMap::from([( - "thinking_effort".to_string(), - serde_json::json!("high"), - )])), - reasoning: Some(false), - ..Default::default() - }; + let mut model_config = + ModelConfig::new("gpt-4").with_merged_request_params(std::collections::HashMap::from( + [("thinking_effort".to_string(), serde_json::json!("high"))], + )); + model_config.reasoning = Some(false); let options = build_config_options( &mode_state, diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index 5b038cf03511..ae2d1e1acdbb 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -104,7 +104,6 @@ mod tools; pub type AcpProviderFactory = Arc< dyn Fn( String, - goose_providers::model::ModelConfig, Vec, Option, ) -> BoxFuture<'static, Result>> @@ -938,17 +937,10 @@ impl GooseAcpAgent { async fn create_provider( &self, provider_name: &str, - model_config: goose_providers::model::ModelConfig, extensions: Vec, working_dir: Option, ) -> Result> { - (self.provider_factory)( - provider_name.to_string(), - model_config, - extensions, - working_dir, - ) - .await + (self.provider_factory)(provider_name.to_string(), extensions, working_dir).await } async fn maybe_refresh_provider_inventory_with_agent( @@ -1468,6 +1460,19 @@ impl GooseAcpAgent { checking network connectivity, listing files in src directory"; let user_text = format!("Tool: {name}\nArguments: {args_json}"); let message = Message::user().with_text(&user_text); + let model_config = match agent.model_config_for_session(&sid.0).await { + Ok(config) => config, + Err(_) => return, + }; + let fast_model_config = match crate::model_config::get_fast_model( + provider.get_name(), + &model_config, + ) + .await + { + Ok(config) => config, + Err(_) => return, + }; // The fast model occasionally returns an empty response // under load (rate limiting, transient network). One // retry with a short backoff is enough to recover the @@ -1475,7 +1480,13 @@ impl GooseAcpAgent { let mut llm_outcome: Option = None; for attempt in 0..2 { match provider - .complete_fast(&sid.0, system, std::slice::from_ref(&message), &[]) + .complete( + &fast_model_config, + &sid.0, + system, + std::slice::from_ref(&message), + &[], + ) .await { Ok((response, _)) => { @@ -1742,6 +1753,16 @@ impl GooseAcpAgent { user_text.push_str(&format!("Step {}: {} {}\n", i + 1, name, args)); } let message = Message::user().with_text(&user_text); + let model_config = match agent.model_config_for_session(&sid.0).await { + Ok(config) => config, + Err(_) => return, + }; + let fast_model_config = + match crate::model_config::get_fast_model(provider.get_name(), &model_config).await + { + Ok(config) => config, + Err(_) => return, + }; // Match the per-tool retry policy: one retry on empty/error keeps // the chain header reliable when the fast model is rate-limited or @@ -1749,7 +1770,13 @@ impl GooseAcpAgent { let mut summary: Option = None; for attempt in 0..2 { match provider - .complete_fast(&sid.0, system, std::slice::from_ref(&message), &[]) + .complete( + &fast_model_config, + &sid.0, + system, + std::slice::from_ref(&message), + &[], + ) .await { Ok((response, _)) => { @@ -2710,7 +2737,10 @@ impl GooseAcpAgent { .await .internal_err_ctx("Failed to get provider")?; let provider_name = current_provider.get_name().to_string(); - let current_model_config = current_provider.get_model_config(); + let current_model_config = agent + .model_config_for_session(session_id) + .await + .internal_err_ctx("Failed to resolve model config")?; let model_config = crate::model_config::model_config_from_user_config_with_session_settings( &provider_name, @@ -2743,7 +2773,10 @@ impl GooseAcpAgent { .await .internal_err_ctx("Failed to get provider")?; let provider_name = provider.get_name().to_string(); - let current_model_config = provider.get_model_config(); + let current_model_config = agent + .model_config_for_session(&session_id.0) + .await + .internal_err_ctx("Failed to resolve model config")?; let current_model = current_model_config.model_name.clone(); let goose_mode = agent.goose_mode().await; let inventory = self @@ -2828,7 +2861,10 @@ impl GooseAcpAgent { .await .internal_err_ctx("Failed to get provider")?; let current_provider_name = current_provider.get_name(); - let current_model_config = current_provider.get_model_config(); + let current_model_config = agent + .model_config_for_session(session_id) + .await + .internal_err_ctx("Failed to resolve model config")?; let current_model = current_model_config.model_name.clone(); let use_default_provider = provider_name == DEFAULT_PROVIDER_ID; let resolved_provider_name = if use_default_provider { @@ -3759,7 +3795,6 @@ print(\"hello, world\") ); session.model_config = Some( goose_providers::model::ModelConfig::new("test-model") - .unwrap() .with_context_limit(Some(258_000)), ); let updates = build_usage_updates(&session).expect("usage updates should be present"); diff --git a/crates/goose/src/acp/server/dispatch.rs b/crates/goose/src/acp/server/dispatch.rs index f01285df8c98..27516c8fe2b7 100644 --- a/crates/goose/src/acp/server/dispatch.rs +++ b/crates/goose/src/acp/server/dispatch.rs @@ -204,7 +204,9 @@ impl HandleDispatchFrom for GooseAcpHandler { .await { Ok(()) => match AssertUnwindSafe( - provider.fetch_recommended_models(), + provider.fetch_recommended_models( + crate::model_config::global_toolshim(), + ), ) .catch_unwind() .await diff --git a/crates/goose/src/acp/server/providers.rs b/crates/goose/src/acp/server/providers.rs index fc5f9671cb1f..342d647d73f3 100644 --- a/crates/goose/src/acp/server/providers.rs +++ b/crates/goose/src/acp/server/providers.rs @@ -443,16 +443,8 @@ impl GooseAcpAgent { &self, req: ProviderSupportedModelsListRequest, ) -> Result { - let entry = crate::providers::get_from_registry(&req.provider_id) - .await - .invalid_params_err_ctx("Unknown provider")?; - let model_config = crate::model_config::model_config_from_user_config( - &req.provider_id, - &entry.metadata().default_model, - ) - .invalid_params_err_ctx("Invalid default model")?; let provider = self - .create_provider(&req.provider_id, model_config, Vec::new(), None) + .create_provider(&req.provider_id, Vec::new(), None) .await .internal_err_ctx("Failed to initialize provider")?; let models = provider @@ -722,35 +714,33 @@ impl GooseAcpAgent { tokio::spawn(async move { let mut refresh_guard = provider_inventory.refresh_guard(&identity); let provider_result = AssertUnwindSafe(async { - let metadata = crate::providers::get_from_registry(&provider_id).await?; - let model_config = crate::model_config::model_config_from_user_config( - &provider_id, - &metadata.metadata().default_model, - )?; - provider_factory(provider_id.clone(), model_config, Vec::new(), None).await + provider_factory(provider_id.clone(), Vec::new(), None).await }) .catch_unwind() .await; - let fetch_result: Result> = match provider_result { - Ok(Ok(provider)) => { - match ensure_refresh_identity_current(&provider_id, &identity).await { - Ok(()) => match AssertUnwindSafe(provider.fetch_recommended_models()) + let fetch_result: Result> = + match provider_result { + Ok(Ok(provider)) => { + match ensure_refresh_identity_current(&provider_id, &identity).await { + Ok(()) => match AssertUnwindSafe(provider.fetch_recommended_models( + crate::model_config::global_toolshim(), + )) .catch_unwind() .await - { - Ok(Ok(models)) => Ok(models), - Ok(Err(error)) => Err(anyhow::anyhow!(error.to_string())), - Err(_) => { - Err(anyhow::anyhow!("provider inventory refresh task panicked")) - } - }, - Err(error) => Err(error), + { + Ok(Ok(models)) => Ok(models), + Ok(Err(error)) => Err(anyhow::anyhow!(error.to_string())), + Err(_) => Err(anyhow::anyhow!( + "provider inventory refresh task panicked" + )), + }, + Err(error) => Err(error), + } } - } - Ok(Err(error)) => Err(error), - Err(_) => Err(anyhow::anyhow!("provider inventory refresh task panicked")), - }; + Ok(Err(error)) => Err(error), + Err(_) => Err(anyhow::anyhow!("provider inventory refresh task panicked")), + }; match fetch_result { Ok(models) => match provider_inventory diff --git a/crates/goose/src/acp/server_factory.rs b/crates/goose/src/acp/server_factory.rs index 7b4149baa2fe..af087a297d5a 100644 --- a/crates/goose/src/acp/server_factory.rs +++ b/crates/goose/src/acp/server_factory.rs @@ -58,26 +58,22 @@ impl AcpServer { let disable_session_naming = config.get_goose_disable_session_naming().unwrap_or(false); let scheduler = self.scheduler().await?; - let provider_factory: AcpProviderFactory = Arc::new( - move |provider_name, model_config, extensions, working_dir| { + let provider_factory: AcpProviderFactory = + Arc::new(move |provider_name, extensions, working_dir| { Box::pin(async move { match working_dir { Some(working_dir) => { crate::providers::create_with_working_dir( &provider_name, - model_config, extensions, working_dir, ) .await } - None => { - crate::providers::create(&provider_name, model_config, extensions).await - } + None => crate::providers::create(&provider_name, extensions).await, } }) - }, - ); + }); let agent = GooseAcpAgent::new(GooseAcpAgentOptions { provider_factory, diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index c2299e83b268..4665d76b3f02 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -138,6 +138,7 @@ pub struct ReplyContext { pub goose_mode: GooseMode, pub tool_call_cut_off: usize, pub initial_messages: Vec, + pub model_config: goose_providers::model::ModelConfig, } pub struct ToolCategorizeResult { @@ -344,6 +345,7 @@ impl Agent { .and_then(|host_info| host_info.client_name.clone()) .unwrap_or_else(|| goose_platform.to_string()); let session_manager = Arc::clone(&config.session_manager); + let inspection_session_manager = Arc::clone(&config.session_manager); let permission_manager = Arc::clone(&config.permission_manager); let use_login_shell_path = config.resolve_use_login_shell_path(); Self { @@ -369,6 +371,7 @@ impl Agent { tool_inspection_manager: Self::create_tool_inspection_manager( permission_manager, provider.clone(), + inspection_session_manager, ), hook_manager: crate::hooks::HookManager::load( std::env::current_dir().ok().as_deref(), @@ -577,6 +580,7 @@ impl Agent { fn create_tool_inspection_manager( permission_manager: Arc, provider: SharedProvider, + session_manager: Arc, ) -> ToolInspectionManager { let mut tool_inspection_manager = ToolInspectionManager::new(); @@ -585,12 +589,16 @@ impl Agent { tool_inspection_manager.add_inspector(Box::new(EgressInspector::new())); // Add adversary inspector (LLM-based review, enabled by ~/.config/goose/adversary.md) - tool_inspection_manager.add_inspector(Box::new(AdversaryInspector::new(provider.clone()))); + tool_inspection_manager.add_inspector(Box::new(AdversaryInspector::new( + provider.clone(), + session_manager.clone(), + ))); // Add permission inspector (medium-high priority) tool_inspection_manager.add_inspector(Box::new(PermissionInspector::new( permission_manager, provider, + session_manager, ))); // Add repetition inspector (lower priority - basic repetition checking) @@ -689,7 +697,7 @@ impl Agent { } let initial_messages = conversation.messages().clone(); - let (tools, toolshim_tools, system_prompt) = self + let (tools, toolshim_tools, system_prompt, model_config) = self .prepare_tools_and_prompt(session_id, working_dir) .await?; @@ -703,11 +711,13 @@ impl Agent { { Ok(v) => v, Err(_) => { - let context_limit = self - .provider() - .await - .map(|p| p.get_model_config().context_limit()) - .unwrap_or(goose_providers::model::DEFAULT_CONTEXT_LIMIT); + let context_limit = match self.provider().await { + Ok(provider) => provider + .get_context_limit(&model_config) + .await + .unwrap_or_else(|_| model_config.context_limit()), + Err(_) => goose_providers::model::DEFAULT_CONTEXT_LIMIT, + }; let compaction_threshold = Config::global() .get_param::("GOOSE_AUTO_COMPACT_THRESHOLD") .unwrap_or(crate::context_mgmt::DEFAULT_COMPACTION_THRESHOLD); @@ -723,6 +733,7 @@ impl Agent { goose_mode, tool_call_cut_off, initial_messages, + model_config, }) } @@ -811,6 +822,38 @@ impl Agent { } } + /// Resolve the active model config for a session. + /// + /// The session is the source of truth for the selected model and its + /// settings. When the session has no stored config (e.g. before the + /// provider has been persisted), fall back to the configured provider + /// defaults. + pub async fn model_config_for_session( + &self, + session_id: &str, + ) -> Result { + if let Ok(session) = self + .config + .session_manager + .get_session(session_id, false) + .await + { + if let Some(model_config) = session.model_config { + return Ok(model_config); + } + } + + let config = Config::global(); + let provider_name = config + .get_goose_provider() + .map_err(|_| anyhow!("Could not resolve model config: missing provider"))?; + let model_name = config + .get_goose_model() + .map_err(|_| anyhow!("Could not resolve model config: missing model"))?; + crate::model_config::model_config_from_user_config(&provider_name, &model_name) + .map_err(|e| anyhow!("Could not resolve model config: {e}")) + } + /// When set, all stdio extensions will be started via `docker exec` in the specified container. pub async fn set_container(&self, container: Option) { *self.container.lock().await = container.clone(); @@ -1671,8 +1714,10 @@ impl Agent { ) ); + let compact_model_config = self.model_config_for_session(&session_config.id).await?; match compact_messages( self.provider().await?.as_ref(), + &compact_model_config, &session_config.id, &conversation_to_compact, false, @@ -1730,6 +1775,7 @@ impl Agent { tool_call_cut_off, goose_mode, initial_messages, + model_config, } = context; if let Some(project_addendum) = self.load_project_instructions(&session).await { @@ -1740,7 +1786,7 @@ impl Agent { let provider = self.provider().await?; let provider_name = provider.get_name().to_string(); - let requested_model = provider.get_model_config().model_name; + let requested_model = model_config.model_name.clone(); let inference = provider .fetch_model_info(&requested_model) .await @@ -1904,6 +1950,7 @@ impl Agent { let mut stream = Self::stream_response_from_provider( self.provider().await?, + model_config.clone(), &session_config.id, &system_prompt, conversation_with_moim.messages(), @@ -1922,6 +1969,7 @@ impl Agent { } else { crate::context_mgmt::maybe_summarize_tool_pairs( self.provider().await?, + model_config.clone(), session_config.id.clone(), conversation.clone(), tool_call_cut_off, @@ -2273,6 +2321,7 @@ impl Agent { match compact_messages( self.provider().await?.as_ref(), + &model_config, &session_config.id, &conversation, false, @@ -2366,7 +2415,7 @@ impl Agent { can_drain_pending_steers = true; if tools_updated { - (tools, toolshim_tools, system_prompt) = + (tools, toolshim_tools, system_prompt, _) = self.prepare_tools_and_prompt(&session_config.id, &session.working_dir).await?; } @@ -2377,7 +2426,7 @@ impl Agent { .await .load_subdirectory_hints(&working_dir); if has_new_hints && !tools_updated { - (tools, toolshim_tools, system_prompt) = + (tools, toolshim_tools, system_prompt, _) = self.prepare_tools_and_prompt(&session_config.id, &session.working_dir).await?; } } @@ -2607,10 +2656,21 @@ impl Agent { pub async fn update_provider( &self, provider: Arc, + model_config: goose_providers::model::ModelConfig, session_id: &str, ) -> Result<()> { let provider_name = provider.get_name().to_string(); - let model_config = provider.get_model_config(); + + // Normalize against the provider entry so custom/declarative providers + // backfill `context_limit` from their known models before the config is + // persisted as the session source of truth; otherwise auto-compaction + // would fall back to DEFAULT_CONTEXT_LIMIT. + let model_config = match crate::providers::get_from_registry(&provider_name).await { + Ok(entry) => entry + .normalize_model_config(model_config.clone()) + .unwrap_or(model_config), + Err(_) => model_config, + }; let mut current_provider = self.provider.lock().await; *current_provider = Some(provider); @@ -2668,14 +2728,14 @@ impl Agent { let provider = crate::providers::create_with_working_dir( provider_name, - model_config, extensions, session.working_dir.clone(), ) .await .map_err(|e| anyhow!("Could not create provider: {}", e))?; - self.update_provider(provider, session_id).await?; + self.update_provider(provider, model_config, session_id) + .await?; let mode = self.goose_mode().await; self.update_goose_mode(mode, session_id).await @@ -2688,8 +2748,9 @@ impl Agent { ) -> Result<()> { let current_provider = self.provider().await?; let provider_name = current_provider.get_name().to_string(); - let model_config = current_provider - .get_model_config() + let model_config = self + .model_config_for_session(session_id) + .await? .with_thinking_effort(effort); self.recreate_provider_for_session(session_id, &provider_name, model_config) @@ -2723,79 +2784,80 @@ impl Agent { let extensions = EnabledExtensionsState::extensions_or_default(Some(&session.extension_data), config); - let (provider, provider_changed) = if crate::providers::get_from_registry(&provider_name) - .await - .is_ok() - { - let p = crate::providers::create_with_working_dir( - &provider_name, - model_config, - extensions, - session.working_dir.clone(), - ) - .await - .map_err(|e| anyhow!("Could not create provider: {}", e))?; - (p, false) - } else { - let fallback_provider_name = config - .get_goose_provider() - .ok() - .filter(|name| name != &provider_name) - .ok_or_else(|| { - anyhow!( - "Could not create provider: provider '{}' not found", - provider_name - ) - })?; - - tracing::warn!( - "Session provider '{}' unavailable, falling back to '{}'", - provider_name, - fallback_provider_name - ); - - let fallback_model_name = config - .get_goose_model() - .ok() - .ok_or_else(|| anyhow!("Could not configure fallback provider: missing model"))?; - let fallback_model_config = crate::model_config::model_config_from_user_config( - &fallback_provider_name, - &fallback_model_name, - ) - .map_err(|e| anyhow!("Could not configure fallback provider: invalid model {}", e))?; + let (provider, active_model_config, provider_changed) = + if crate::providers::get_from_registry(&provider_name) + .await + .is_ok() + { + let p = crate::providers::create_with_working_dir( + &provider_name, + extensions, + session.working_dir.clone(), + ) + .await + .map_err(|e| anyhow!("Could not create provider: {}", e))?; + (p, model_config, false) + } else { + let fallback_provider_name = config + .get_goose_provider() + .ok() + .filter(|name| name != &provider_name) + .ok_or_else(|| { + anyhow!( + "Could not create provider: provider '{}' not found", + provider_name + ) + })?; - let fallback_provider = crate::providers::create_with_working_dir( - &fallback_provider_name, - fallback_model_config.clone(), - extensions, - session.working_dir.clone(), - ) - .await - .map_err(|e| { - anyhow!( - "Could not create provider '{}' or fallback '{}': {}", + tracing::warn!( + "Session provider '{}' unavailable, falling back to '{}'", provider_name, - fallback_provider_name, - e + fallback_provider_name + ); + + let fallback_model_name = config.get_goose_model().ok().ok_or_else(|| { + anyhow!("Could not configure fallback provider: missing model") + })?; + let fallback_model_config = crate::model_config::model_config_from_user_config( + &fallback_provider_name, + &fallback_model_name, ) - })?; + .map_err(|e| { + anyhow!("Could not configure fallback provider: invalid model {}", e) + })?; - if let Err(e) = self - .config - .session_manager - .update(&session.id) - .provider_name(&fallback_provider_name) - .model_config(fallback_model_config) - .apply() + let fallback_provider = crate::providers::create_with_working_dir( + &fallback_provider_name, + extensions, + session.working_dir.clone(), + ) .await - { - tracing::warn!("Failed to update session provider: {}", e); - } + .map_err(|e| { + anyhow!( + "Could not create provider '{}' or fallback '{}': {}", + provider_name, + fallback_provider_name, + e + ) + })?; - (fallback_provider, true) - }; + if let Err(e) = self + .config + .session_manager + .update(&session.id) + .provider_name(&fallback_provider_name) + .model_config(fallback_model_config.clone()) + .apply() + .await + { + tracing::warn!("Failed to update session provider: {}", e); + } - self.update_provider(provider, &session.id).await?; + (fallback_provider, fallback_model_config, true) + }; + + self.update_provider(provider, active_model_config, &session.id) + .await?; // Propagate session mode to the new provider if let Some(provider) = self.provider.lock().await.as_ref() { provider @@ -2909,12 +2971,7 @@ impl Agent { tracing::debug!("Retrieved {} extensions info", extensions_info.len()); let (extension_count, tool_count) = self.total_extension_and_tool_counts(session_id).await; - // Get model name from provider - let provider = self.provider().await.map_err(|e| { - tracing::error!("Failed to get provider for recipe creation: {}", e); - e - })?; - let model_config = provider.get_model_config(); + let model_config = self.model_config_for_session(session_id).await?; let model_name = &model_config.model_name; tracing::debug!("Using model: {}", model_name); @@ -2956,15 +3013,6 @@ impl Agent { ); tracing::info!("Calling provider to generate recipe content"); - let model_config = { - let provider_guard = self.provider.lock().await; - let provider = provider_guard.as_ref().ok_or_else(|| { - let error = anyhow!("Provider not available during recipe creation"); - tracing::error!("{}", error); - error - })?; - provider.get_model_config() - }; let (result, _usage) = self .provider .lock() @@ -3191,9 +3239,6 @@ mod tests { fn get_name(&self) -> &str { "test-action-required" } - fn get_model_config(&self) -> goose_providers::model::ModelConfig { - goose_providers::model::ModelConfig::new("test").unwrap() - } async fn stream( &self, _: &goose_providers::model::ModelConfig, @@ -3390,10 +3435,6 @@ exit 0 Ok(stream_from_single_message(message, usage)) } - fn get_model_config(&self) -> goose_providers::model::ModelConfig { - goose_providers::model::ModelConfig::new("mock-model").unwrap() - } - fn get_name(&self) -> &str { "counting-text" } @@ -3422,10 +3463,6 @@ exit 0 }))) } - fn get_model_config(&self) -> goose_providers::model::ModelConfig { - goose_providers::model::ModelConfig::new("mock-model").unwrap() - } - fn get_name(&self) -> &str { "refusing" } @@ -3497,7 +3534,13 @@ exit 0 GooseMode::Auto, ) .await?; - agent.update_provider(provider, &session.id).await?; + agent + .update_provider( + provider, + goose_providers::model::ModelConfig::new("mock-model"), + &session.id, + ) + .await?; Ok((agent, session.id)) } diff --git a/crates/goose/src/agents/execute_commands.rs b/crates/goose/src/agents/execute_commands.rs index c82f459ba0f6..84e9152b0382 100644 --- a/crates/goose/src/agents/execute_commands.rs +++ b/crates/goose/src/agents/execute_commands.rs @@ -156,8 +156,10 @@ impl Agent { .conversation .ok_or_else(|| anyhow!("Session has no conversation"))?; + let model_config = self.model_config_for_session(session_id).await?; let (compacted_conversation, usage) = compact_messages( self.provider().await?.as_ref(), + &model_config, session_id, &conversation, true, // is_manual_compact @@ -209,8 +211,11 @@ impl Agent { async fn handle_status_command(&self, session_id: &str) -> Result> { let provider = self.provider().await?; - let model_config = provider.get_model_config(); - let context_limit = model_config.context_limit(); + let model_config = self.model_config_for_session(session_id).await?; + let context_limit = provider + .get_context_limit(&model_config) + .await + .unwrap_or_else(|_| model_config.context_limit()); let goose_mode = self.goose_mode().await; diff --git a/crates/goose/src/agents/mcp_client.rs b/crates/goose/src/agents/mcp_client.rs index 56567a005065..913a7a816fe0 100644 --- a/crates/goose/src/agents/mcp_client.rs +++ b/crates/goose/src/agents/mcp_client.rs @@ -40,6 +40,13 @@ pub type Error = rmcp::ServiceError; const MCP_APPS_UI_EXTENSION_ID: &str = "io.modelcontextprotocol/ui"; const MCP_APPS_UI_MIME_TYPE: &str = "text/html;profile=mcp-app"; +fn resolve_sampling_model_config() -> anyhow::Result { + let config = crate::config::Config::global(); + let provider_name = config.get_goose_provider()?; + let model_name = config.get_goose_model()?; + crate::model_config::model_config_from_user_config(&provider_name, &model_name) +} + fn default_mcp_apps_ui_extensions() -> ExtensionCapabilities { let mut extensions = ExtensionCapabilities::new(); let mut ui_extension_settings = JsonObject::new(); @@ -324,7 +331,13 @@ impl ClientHandler for GooseClient { .as_deref() .unwrap_or("You are a general-purpose AI agent called goose"); - let model_config = provider.get_model_config(); + let model_config = resolve_sampling_model_config().map_err(|e| { + ErrorData::new( + ErrorCode::INTERNAL_ERROR, + "Could not resolve model config", + Some(Value::from(e.to_string())), + ) + })?; let (response, usage) = provider .complete( &model_config, diff --git a/crates/goose/src/agents/moim.rs b/crates/goose/src/agents/moim.rs index 5bc8a79b15be..a6b8407a32b9 100644 --- a/crates/goose/src/agents/moim.rs +++ b/crates/goose/src/agents/moim.rs @@ -51,23 +51,22 @@ pub async fn inject_moim( .get_session(session_id, false) .await .ok(); - let provider_context_limit = - extension_manager - .get_provider() - .try_lock() - .ok() - .and_then(|provider| { - provider - .as_ref() - .map(|provider| provider.get_model_config().context_limit()) - }); - let session_context_limit = session.as_ref().and_then(|session| { - session - .model_config - .as_ref() - .map(|config| config.context_limit()) - }); - let context_limit = provider_context_limit.or(session_context_limit); + let session_model_config = session + .as_ref() + .and_then(|session| session.model_config.clone()); + let context_limit = if let Some(model_config) = session_model_config.as_ref() { + let provider = extension_manager.get_provider().lock().await.clone(); + match provider { + Some(provider) => provider + .get_context_limit(model_config) + .await + .ok() + .or_else(|| Some(model_config.context_limit())), + None => Some(model_config.context_limit()), + } + } else { + None + }; if should_skip_moim(context_limit) { return conversation; } diff --git a/crates/goose/src/agents/platform_extensions/apps.rs b/crates/goose/src/agents/platform_extensions/apps.rs index e3e4644a9c52..ee18b586df4c 100644 --- a/crates/goose/src/agents/platform_extensions/apps.rs +++ b/crates/goose/src/agents/platform_extensions/apps.rs @@ -272,7 +272,7 @@ impl AppsManagerClient { let messages = vec![Message::user().with_text(&user_prompt)]; let tools = vec![Self::create_app_content_tool()]; - let model_config = provider.get_model_config(); + let model_config = self.context.model_config_for_session(session_id).await?; let (response, usage) = provider .complete(&model_config, session_id, &system_prompt, &messages, &tools) @@ -315,7 +315,7 @@ impl AppsManagerClient { let messages = vec![Message::user().with_text(&user_prompt)]; let tools = vec![Self::update_app_content_tool()]; - let model_config = provider.get_model_config(); + let model_config = self.context.model_config_for_session(session_id).await?; let (response, usage) = provider .complete(&model_config, session_id, &system_prompt, &messages, &tools) diff --git a/crates/goose/src/agents/platform_extensions/mod.rs b/crates/goose/src/agents/platform_extensions/mod.rs index 4b404937d016..adef19cf99eb 100644 --- a/crates/goose/src/agents/platform_extensions/mod.rs +++ b/crates/goose/src/agents/platform_extensions/mod.rs @@ -216,6 +216,27 @@ pub struct PlatformExtensionContext { } impl PlatformExtensionContext { + pub async fn model_config_for_session( + &self, + session_id: &str, + ) -> Result { + if let Ok(session) = self.session_manager.get_session(session_id, false).await { + if let Some(model_config) = session.model_config { + return Ok(model_config); + } + } + + let config = crate::config::Config::global(); + let provider_name = config + .get_goose_provider() + .map_err(|_| "Could not resolve model config: missing provider".to_string())?; + let model_name = config + .get_goose_model() + .map_err(|_| "Could not resolve model config: missing model".to_string())?; + crate::model_config::model_config_from_user_config(&provider_name, &model_name) + .map_err(|e| format!("Could not resolve model config: {e}")) + } + pub fn result_with_platform_notification( &self, mut result: rmcp::model::CallToolResult, diff --git a/crates/goose/src/agents/platform_extensions/orchestrator.rs b/crates/goose/src/agents/platform_extensions/orchestrator.rs index 458003435802..c56827b9b088 100644 --- a/crates/goose/src/agents/platform_extensions/orchestrator.rs +++ b/crates/goose/src/agents/platform_extensions/orchestrator.rs @@ -141,6 +141,21 @@ impl OrchestratorClient { .ok_or_else(|| "Provider not available".to_string()) } + async fn parent_model_config( + &self, + provider_name: &str, + ) -> Result { + if let Some(session) = self.context.session.as_ref() { + return self.context.model_config_for_session(&session.id).await; + } + + let model_name = Config::global() + .get_goose_model() + .map_err(|_| "Could not resolve model config: missing model".to_string())?; + crate::model_config::model_config_from_user_config(provider_name, &model_name) + .map_err(|e| format!("Could not resolve model config: {e}")) + } + fn parent_extensions(&self) -> Vec { let extension_data = self.context.session.as_ref().map(|s| &s.extension_data); EnabledExtensionsState::extensions_or_default(extension_data, Config::global()) @@ -337,10 +352,17 @@ impl OrchestratorClient { conversation_text )); - let (response, _usage) = provider - .complete_fast(session_id, system, &[user_message], &[]) - .await - .map_err(|e| format!("LLM summarization failed: {}", e))?; + let model_config = self.parent_model_config(provider.get_name()).await?; + let (response, _usage) = crate::model_config::complete_fast( + provider.as_ref(), + &model_config, + session_id, + system, + &[user_message], + &[], + ) + .await + .map_err(|e| format!("LLM summarization failed: {}", e))?; Ok(response .content @@ -406,15 +428,12 @@ impl OrchestratorClient { let parent_provider = self.get_provider().await?; let extensions = self.parent_extensions(); - let provider = providers::create( - parent_provider.get_name(), - parent_provider.get_model_config(), - extensions, - ) - .await - .map_err(|e| format!("Failed to create provider for new agent: {}", e))?; + let model_config = self.parent_model_config(parent_provider.get_name()).await?; + let provider = providers::create(parent_provider.get_name(), extensions) + .await + .map_err(|e| format!("Failed to create provider for new agent: {}", e))?; agent - .update_provider(provider, &session.id) + .update_provider(provider, model_config, &session.id) .await .map_err(|e| format!("Failed to set provider on new agent: {}", e))?; @@ -448,15 +467,12 @@ impl OrchestratorClient { if agent.provider().await.is_err() { if let Ok(parent_provider) = self.get_provider().await { let extensions = self.parent_extensions(); - if let Ok(provider) = providers::create( - parent_provider.get_name(), - parent_provider.get_model_config(), - extensions, - ) - .await + let model_config = self.parent_model_config(parent_provider.get_name()).await?; + if let Ok(provider) = + providers::create(parent_provider.get_name(), extensions).await { agent - .update_provider(provider, &session_id) + .update_provider(provider, model_config, &session_id) .await .map_err(|e| format!("Failed to set provider: {}", e))?; } diff --git a/crates/goose/src/agents/platform_extensions/summarize.rs b/crates/goose/src/agents/platform_extensions/summarize.rs index ea974f647a3e..fb405b93a46d 100644 --- a/crates/goose/src/agents/platform_extensions/summarize.rs +++ b/crates/goose/src/agents/platform_extensions/summarize.rs @@ -149,7 +149,16 @@ impl McpClientTrait for SummarizeClient { }; let session_id = &ctx.session_id; - match execute_summarize(provider, session_id, params, &working_dir).await { + let model_config = match self.context.model_config_for_session(session_id).await { + Ok(config) => config, + Err(e) => { + return Ok(CallToolResult::error(vec![Content::text(format!( + "Error: {}", + e + ))])); + } + }; + match execute_summarize(provider, model_config, session_id, params, &working_dir).await { Ok(result) => Ok(result), Err(msg) => Ok(CallToolResult::error(vec![Content::text(format!( "Error: {}", @@ -165,6 +174,7 @@ impl McpClientTrait for SummarizeClient { async fn execute_summarize( provider: Arc, + model_config: goose_providers::model::ModelConfig, session_id: &str, params: SummarizeParams, working_dir: &Path, @@ -187,8 +197,6 @@ async fn execute_summarize( let user_message = Message::user().with_text(&prompt); - let model_config = provider.get_model_config(); - let (response, _usage) = provider .complete(&model_config, session_id, system, &[user_message], &[]) .await diff --git a/crates/goose/src/agents/platform_extensions/summon.rs b/crates/goose/src/agents/platform_extensions/summon.rs index af22aeef5968..0e128cf4c17b 100644 --- a/crates/goose/src/agents/platform_extensions/summon.rs +++ b/crates/goose/src/agents/platform_extensions/summon.rs @@ -1528,7 +1528,7 @@ impl SummonClient { recipe: &Recipe, session: &crate::session::Session, ) -> Result { - let provider = self.resolve_provider(params, recipe, session).await?; + let (provider, model_config) = self.resolve_provider(params, recipe, session).await?; let mut extensions = EnabledExtensionsState::extensions_or_default( Some(&session.extension_data), @@ -1574,9 +1574,14 @@ impl SummonClient { None => session.working_dir.clone(), }; - let task_config = - TaskConfig::new(provider, &session.id, &effective_working_dir, extensions) - .with_max_turns(Some(max_turns)); + let task_config = TaskConfig::new( + provider, + model_config, + &session.id, + &effective_working_dir, + extensions, + ) + .with_max_turns(Some(max_turns)); Ok(task_config) } @@ -1614,7 +1619,6 @@ impl SummonClient { crate::model_config::model_config_from_user_config(provider_name, &model)?; cfg.toolshim = parent.toolshim; cfg.toolshim_model = parent.toolshim_model; - cfg.fast_model_config = parent.fast_model_config; cfg.temperature = cfg.temperature.or(parent.temperature); if let Some(parent_params) = parent.request_params { let merged = cfg.request_params.get_or_insert_with(Default::default); @@ -1640,7 +1644,13 @@ impl SummonClient { params: &DelegateParams, recipe: &Recipe, session: &crate::session::Session, - ) -> Result, anyhow::Error> { + ) -> Result< + ( + Arc, + goose_providers::model::ModelConfig, + ), + anyhow::Error, + > { let provider_name = params .provider .clone() @@ -1659,7 +1669,8 @@ impl SummonClient { .ok_or_else(|| anyhow::anyhow!("No provider configured"))?; let model_config = self.resolve_model_config(params, recipe, session, &provider_name)?; - providers::create(&provider_name, model_config, Vec::new()).await + let provider = providers::create(&provider_name, Vec::new()).await?; + Ok((provider, model_config)) } fn resolve_max_turns(&self, session: &crate::session::Session) -> usize { @@ -2577,9 +2588,7 @@ You review code."#; } fn parent_config() -> goose_providers::model::ModelConfig { - goose_providers::model::ModelConfig::new(PARENT_MODEL) - .unwrap() - .with_canonical_limits(PROVIDER) + goose_providers::model::ModelConfig::new(PARENT_MODEL).with_canonical_limits(PROVIDER) } #[tokio::test] @@ -2593,7 +2602,6 @@ You review code."#; let parent = parent_config(); let overridden = goose_providers::model::ModelConfig::new(OVERRIDE_MODEL) - .unwrap() .with_canonical_limits(PROVIDER); assert_ne!(parent.context_limit, overridden.context_limit); assert_ne!(parent.reasoning, overridden.reasoning); diff --git a/crates/goose/src/agents/reply_parts.rs b/crates/goose/src/agents/reply_parts.rs index 7d34581c64f2..3e7f6f816b60 100644 --- a/crates/goose/src/agents/reply_parts.rs +++ b/crates/goose/src/agents/reply_parts.rs @@ -22,10 +22,15 @@ use crate::providers::toolshim::{ modify_system_prompt_for_tool_json, sanitize_residual_markers, }; use goose_providers::conversation::token_usage::{ProviderUsage, Usage}; +use goose_providers::model::ModelConfig; use rmcp::model::Tool; use tracing::warn; -async fn enhance_model_error(error: ProviderError, provider: &Arc) -> ProviderError { +async fn enhance_model_error( + error: ProviderError, + provider: &Arc, + toolshim: bool, +) -> ProviderError { let ProviderError::RequestFailed(ref msg) = error else { return error; }; @@ -35,7 +40,7 @@ async fn enhance_model_error(error: ProviderError, provider: &Arc) return error; } - let Ok(models) = provider.fetch_recommended_models().await else { + let Ok(models) = provider.fetch_recommended_models(toolshim).await else { return error; }; if models.is_empty() { @@ -143,7 +148,7 @@ impl Agent { &self, session_id: &str, working_dir: &std::path::Path, - ) -> Result<(Vec, Vec, String)> { + ) -> Result<(Vec, Vec, String, ModelConfig)> { let mut tools = self.list_tools(session_id, None).await; #[cfg(feature = "code-mode")] @@ -215,9 +220,7 @@ impl Agent { .await; let (extension_count, tool_count) = self.total_extension_and_tool_counts(session_id).await; - // Get model name from provider - let provider = self.provider().await?; - let model_config = provider.get_model_config(); + let model_config = self.model_config_for_session(session_id).await?; let goose_mode = *self.current_goose_mode.lock().await; @@ -243,22 +246,23 @@ impl Agent { tools = vec![]; } - Ok((tools, toolshim_tools, system_prompt)) + Ok((tools, toolshim_tools, system_prompt, model_config)) } #[tracing::instrument( - skip(provider, session_id, system_prompt, messages, tools, toolshim_tools), + skip(provider, model_config, session_id, system_prompt, messages, tools, toolshim_tools), fields(session.id = %session_id) )] pub(crate) async fn stream_response_from_provider( provider: Arc, + model_config: ModelConfig, session_id: &str, system_prompt: &str, messages: &[Message], tools: &[Tool], toolshim_tools: &[Tool], ) -> Result { - let config = provider.get_model_config(); + let config = model_config.clone(); let filtered_messages: Vec = messages .iter() @@ -281,9 +285,8 @@ impl Agent { // Capture errors during stream creation and return them as part of the stream // so they can be handled by the existing error handling logic in the agent - let model_config = provider - .get_model_config() - .with_default_thinking_effort(Config::global().get_goose_thinking_effort()); + let model_config = + model_config.with_default_thinking_effort(Config::global().get_goose_thinking_effort()); debug!("WAITING_LLM_STREAM_START"); let stream_result = provider .stream( @@ -300,7 +303,7 @@ impl Agent { let mut stream = match stream_result { Ok(s) => s, Err(e) => { - let enhanced_error = enhance_model_error(e, &provider).await; + let enhanced_error = enhance_model_error(e, &provider, config.toolshim).await; // Return a stream that immediately yields the error // This allows the error to be caught by existing error handling in agent.rs return Ok(Box::pin(try_stream! { @@ -621,9 +624,7 @@ mod tests { use rmcp::object; #[derive(Clone)] - struct MockProvider { - model_config: ModelConfig, - } + struct MockProvider; #[async_trait] impl Provider for MockProvider { @@ -631,10 +632,6 @@ mod tests { "mock" } - fn get_model_config(&self) -> ModelConfig { - self.model_config.clone() - } - async fn stream( &self, _model_config: &ModelConfig, @@ -664,9 +661,11 @@ mod tests { ) .await?; - let model_config = ModelConfig::new("test-model").unwrap(); - let provider = std::sync::Arc::new(MockProvider { model_config }); - agent.update_provider(provider, &session.id).await?; + let model_config = ModelConfig::new("test-model"); + let provider = std::sync::Arc::new(MockProvider); + agent + .update_provider(provider, model_config, &session.id) + .await?; // Add unsorted frontend tools let frontend_tools = vec![ @@ -697,7 +696,7 @@ mod tests { .await .unwrap(); - let (tools, _toolshim_tools, _system_prompt) = agent + let (tools, _toolshim_tools, _system_prompt, _model_config) = agent .prepare_tools_and_prompt(&session.id, session.working_dir.as_path()) .await?; diff --git a/crates/goose/src/agents/subagent_handler.rs b/crates/goose/src/agents/subagent_handler.rs index b506376ef793..6c37debbb9b7 100644 --- a/crates/goose/src/agents/subagent_handler.rs +++ b/crates/goose/src/agents/subagent_handler.rs @@ -141,7 +141,11 @@ fn get_agent_messages(params: SubagentRunParams) -> AgentMessagesFuture { let agent = Arc::new(Agent::with_config(config)); agent - .update_provider(task_config.provider.clone(), &session_id) + .update_provider( + task_config.provider.clone(), + task_config.model_config.clone(), + &session_id, + ) .await .map_err(|e| anyhow!("Failed to set provider on sub agent: {}", e))?; diff --git a/crates/goose/src/agents/subagent_task_config.rs b/crates/goose/src/agents/subagent_task_config.rs index 1cca9c086cbd..33e99ab01c2d 100644 --- a/crates/goose/src/agents/subagent_task_config.rs +++ b/crates/goose/src/agents/subagent_task_config.rs @@ -12,6 +12,7 @@ pub const DEFAULT_SUBAGENT_MAX_TURNS: usize = 25; #[derive(Clone)] pub struct TaskConfig { pub provider: Arc, + pub model_config: goose_providers::model::ModelConfig, pub parent_session_id: String, pub parent_working_dir: PathBuf, pub extensions: Vec, @@ -33,12 +34,14 @@ impl fmt::Debug for TaskConfig { impl TaskConfig { pub fn new( provider: Arc, + model_config: goose_providers::model::ModelConfig, parent_session_id: &str, parent_working_dir: &Path, extensions: Vec, ) -> Self { Self { provider, + model_config, parent_session_id: parent_session_id.to_owned(), parent_working_dir: parent_working_dir.to_owned(), extensions, diff --git a/crates/goose/src/bin/build_canonical_models.rs b/crates/goose/src/bin/build_canonical_models.rs index e57d9c2fe60f..cc9a25a84935 100644 --- a/crates/goose/src/bin/build_canonical_models.rs +++ b/crates/goose/src/bin/build_canonical_models.rs @@ -619,11 +619,11 @@ async fn build_canonical_models() -> Result<()> { async fn check_provider( provider_name: &str, - model_for_init: &str, + _model_for_init: &str, ) -> Result<(Vec, Vec, Vec)> { println!("Checking provider: {}", provider_name); - let provider = match create_with_named_model(provider_name, model_for_init, Vec::new()).await { + let provider = match create_with_named_model(provider_name, Vec::new()).await { Ok(p) => p, Err(e) => { println!(" ⚠ Failed to create provider: {}", e); @@ -644,7 +644,7 @@ async fn check_provider( } }; - let recommended_models = match provider.fetch_recommended_models().await { + let recommended_models = match provider.fetch_recommended_models(false).await { Ok(models) => { println!(" ✓ Found {} recommended models", models.len()); models diff --git a/crates/goose/src/config/declarative_providers.rs b/crates/goose/src/config/declarative_providers.rs index 3ccc3aeb7f19..fbac06f066c9 100644 --- a/crates/goose/src/config/declarative_providers.rs +++ b/crates/goose/src/config/declarative_providers.rs @@ -585,10 +585,10 @@ pub fn register_declarative_provider( &config, provider_type, config.dynamic_models.unwrap_or(false), - move |model, tls_config| { + move |tls_config| { let mut cfg = captured.clone(); resolve_config(&mut cfg)?; - HuggingFaceProvider::from_custom_config(model, cfg, tls_config) + HuggingFaceProvider::from_custom_config(cfg, tls_config) }, move || { let mut cfg = identity_config.clone(); @@ -608,10 +608,10 @@ pub fn register_declarative_provider( &config, provider_type, config.dynamic_models.unwrap_or(false), - move |model, tls_config| { + move |tls_config| { let mut cfg = captured.clone(); resolve_config(&mut cfg)?; - crate::providers::openai_def::from_custom_config(model, cfg, tls_config) + crate::providers::openai_def::from_custom_config(cfg, tls_config) }, move || { let mut cfg = identity_config.clone(); @@ -628,10 +628,10 @@ pub fn register_declarative_provider( &config, provider_type, config.dynamic_models.unwrap_or(false), - move |model, tls_config| { + move |tls_config| { let mut cfg = captured.clone(); resolve_config(&mut cfg)?; - OllamaProvider::from_custom_config(model, cfg, tls_config) + OllamaProvider::from_custom_config(cfg, tls_config) }, move || { let mut cfg = identity_config.clone(); @@ -647,10 +647,10 @@ pub fn register_declarative_provider( &config, provider_type, config.dynamic_models.unwrap_or(false), - move |model, tls_config| { + move |tls_config| { let mut cfg = captured.clone(); resolve_config(&mut cfg)?; - AnthropicProvider::from_custom_config(model, cfg, tls_config) + AnthropicProvider::from_custom_config(cfg, tls_config) }, move || { let mut cfg = identity_config.clone(); diff --git a/crates/goose/src/context_mgmt/mod.rs b/crates/goose/src/context_mgmt/mod.rs index 9db9b40e0eac..3553d4fc66cc 100644 --- a/crates/goose/src/context_mgmt/mod.rs +++ b/crates/goose/src/context_mgmt/mod.rs @@ -9,6 +9,7 @@ use crate::{config::Config, token_counter::create_token_counter}; use anyhow::Result; use goose_providers::conversation::token_usage::ProviderUsage; use goose_providers::errors::ProviderError; +use goose_providers::model::ModelConfig; use indoc::indoc; use rmcp::model::Role; use serde::Serialize; @@ -65,6 +66,7 @@ struct SummarizeContext { /// - `ProviderUsage`: Provider usage from summarization pub async fn compact_messages( provider: &dyn Provider, + model_config: &ModelConfig, session_id: &str, conversation: &Conversation, manual_compact: bool, @@ -128,7 +130,7 @@ pub async fn compact_messages( let messages_to_compact = messages.as_slice(); let (summary_message, summarization_usage) = - do_compact(provider, session_id, messages_to_compact).await?; + do_compact(provider, model_config, session_id, messages_to_compact).await?; // Create the final message list with updated visibility metadata: // 1. Original messages become user_visible but not agent_visible @@ -201,7 +203,14 @@ pub async fn check_if_compaction_needed( .unwrap_or(DEFAULT_COMPACTION_THRESHOLD) }); - let context_limit = provider.get_model_config().context_limit(); + let model_config = session + .model_config + .clone() + .unwrap_or_else(|| ModelConfig::new("unknown")); + let context_limit = provider + .get_context_limit(&model_config) + .await + .unwrap_or_else(|_| model_config.context_limit()); let (current_tokens, _token_source) = match session.usage.total_tokens { Some(tokens) => (tokens as usize, "session metadata"), @@ -282,6 +291,7 @@ fn filter_tool_responses(messages: &[Message], remove_percent: u32) -> Vec<&Mess async fn do_compact( provider: &dyn Provider, + model_config: &ModelConfig, session_id: &str, messages: &[Message], ) -> Result<(Message, ProviderUsage), anyhow::Error> { @@ -313,9 +323,15 @@ async fn do_compact( .with_text("Please summarize the conversation history provided in the system prompt."); let summarization_request = vec![user_message]; - match provider - .complete_fast(session_id, &system_prompt, &summarization_request, &[]) - .await + match crate::model_config::complete_fast( + provider, + model_config, + session_id, + &system_prompt, + &summarization_request, + &[], + ) + .await { Ok((mut response, mut provider_usage)) => { response.role = Role::User; @@ -476,6 +492,7 @@ pub fn tool_ids_to_summarize( pub async fn summarize_tool_call( provider: &dyn Provider, + model_config: &ModelConfig, session_id: &str, conversation: &Conversation, tool_id: &str, @@ -522,9 +539,15 @@ pub async fn summarize_tool_call( if that is what it was. "#}; - let (mut response, _) = provider - .complete_fast(session_id, system_prompt, &summarization_request, &[]) - .await?; + let (mut response, _) = crate::model_config::complete_fast( + provider, + model_config, + session_id, + system_prompt, + &summarization_request, + &[], + ) + .await?; response.role = Role::User; response.created = matching_messages.last().unwrap().created; @@ -535,6 +558,7 @@ pub async fn summarize_tool_call( pub fn maybe_summarize_tool_pairs( provider: Arc, + model_config: ModelConfig, session_id: String, conversation: Conversation, cutoff: usize, @@ -552,7 +576,14 @@ pub fn maybe_summarize_tool_pairs( Some(tokio::spawn(async move { let mut results = Vec::new(); for tool_id in tool_ids { - match summarize_tool_call(provider.as_ref(), &session_id, &conversation, &tool_id).await + match summarize_tool_call( + provider.as_ref(), + &model_config, + &session_id, + &conversation, + &tool_id, + ) + .await { Ok(summary) => results.push((summary, tool_id)), Err(e) => { @@ -569,8 +600,6 @@ mod tests { use super::*; use async_trait::async_trait; use goose_providers::conversation::token_usage::Usage; - use goose_providers::errors::ProviderError; - use goose_providers::model::ModelConfig; use rmcp::model::{AnnotateAble, CallToolRequestParams, RawContent, Tool}; fn create_tool_pair( @@ -614,7 +643,6 @@ mod tests { max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }, @@ -666,8 +694,11 @@ mod tests { Ok(stream_from_single_message(message, usage)) } - fn get_model_config(&self) -> ModelConfig { - self.config.clone() + async fn get_context_limit( + &self, + _model_config: &ModelConfig, + ) -> Result { + Ok(self.config.context_limit()) } } @@ -688,10 +719,16 @@ mod tests { ]; let conversation = Conversation::new_unvalidated(basic_conversation); - let (compacted_conversation, _usage) = - compact_messages(&provider, "test-session-id", &conversation, false) - .await - .unwrap(); + let model_config = provider.config.clone(); + let (compacted_conversation, _usage) = compact_messages( + &provider, + &model_config, + "test-session-id", + &conversation, + false, + ) + .await + .unwrap(); let agent_conversation = compacted_conversation.agent_visible_messages(); @@ -721,7 +758,15 @@ mod tests { } let conversation = Conversation::new_unvalidated(messages); - let result = compact_messages(&provider, "test-session-id", &conversation, false).await; + let model_config = provider.config.clone(); + let result = compact_messages( + &provider, + &model_config, + "test-session-id", + &conversation, + false, + ) + .await; assert!( result.is_ok(), diff --git a/crates/goose/src/doctor.rs b/crates/goose/src/doctor.rs index 5d9624682c68..88a13528d02b 100644 --- a/crates/goose/src/doctor.rs +++ b/crates/goose/src/doctor.rs @@ -78,9 +78,9 @@ async fn ensure_working_provider( } log.push(format!("Looking for alternative models on {} ...", pname)); - if let Some(working) = try_other_models(pname, mname, &mut log).await { - let new_model = working.get_model_config().model_name.clone(); - save_and_set(agent, session_id, working).await?; + if let Some((working, model_config)) = try_other_models(pname, mname, &mut log).await { + let new_model = model_config.model_name.clone(); + save_and_set(agent, session_id, working, model_config).await?; let preamble = log.join("\n"); return Ok(Some(Message::assistant().with_text(format!( "**Goose Doctor**\n\n{}\n\n\ @@ -95,10 +95,10 @@ async fn ensure_working_provider( log.push("Looking for other configured providers ...".to_string()); let skip = provider_name.as_deref().unwrap_or(""); - if let Some(working) = try_other_providers(skip, &mut log).await { + if let Some((working, model_config)) = try_other_providers(skip, &mut log).await { let name = working.get_name().to_string(); - let model = working.get_model_config().model_name.clone(); - save_and_set(agent, session_id, working).await?; + let model = model_config.model_name.clone(); + save_and_set(agent, session_id, working, model_config).await?; let preamble = log.join("\n"); return Ok(Some(Message::assistant().with_text(format!( "**Goose Doctor**\n\n{}\n\n\ @@ -139,21 +139,23 @@ async fn save_and_set( agent: &crate::agents::Agent, session_id: &str, provider: Arc, + model_config: goose_providers::model::ModelConfig, ) -> anyhow::Result<()> { let config = Config::global(); - crate::config::set_active_provider( - config, - provider.get_name(), - &provider.get_model_config().model_name, - )?; - agent.update_provider(provider, session_id).await + crate::config::set_active_provider(config, provider.get_name(), &model_config.model_name)?; + agent + .update_provider(provider, model_config, session_id) + .await } -async fn test_provider(provider: &dyn Provider) -> Result<(), ProviderError> { +async fn test_provider( + provider: &dyn Provider, + model_config: &goose_providers::model::ModelConfig, +) -> Result<(), ProviderError> { let messages = vec![Message::user().with_text("Say 'hello' and nothing else.")]; provider .complete( - &provider.get_model_config(), + model_config, "doctor-check", "Respond as briefly as possible.", &messages, @@ -166,27 +168,30 @@ async fn test_provider(provider: &dyn Provider) -> Result<(), ProviderError> { async fn try_create_and_test( provider_name: &str, model_name: &str, -) -> Result, ProviderError> { +) -> Result<(Arc, goose_providers::model::ModelConfig), ProviderError> { let model_config = crate::model_config::model_config_from_user_config(provider_name, model_name) .map_err(|e| ProviderError::ExecutionError(e.to_string()))?; - let provider = providers::create(provider_name, model_config, vec![]) + let provider = providers::create(provider_name, vec![]) .await .map_err(|e| ProviderError::ExecutionError(e.to_string()))?; - test_provider(provider.as_ref()).await?; - Ok(provider) + test_provider(provider.as_ref(), &model_config).await?; + Ok((provider, model_config)) } async fn try_other_models( provider_name: &str, skip_model: &str, log: &mut Vec, -) -> Option> { +) -> Option<(Arc, goose_providers::model::ModelConfig)> { let entry = providers::get_from_registry(provider_name).await.ok()?; let temp = entry.create_with_default_model(vec![]).await.ok()?; - let models = temp.fetch_recommended_models().await.ok()?; + let toolshim = Config::global() + .get_param::("GOOSE_TOOLSHIM") + .unwrap_or(false); + let models = temp.fetch_recommended_models(toolshim).await.ok()?; for model in models.iter().filter(|m| m.as_str() != skip_model).take(3) { log.push(format!(" Trying {} / {} ...", provider_name, model)); @@ -201,7 +206,10 @@ async fn try_other_models( None } -async fn try_other_providers(skip: &str, log: &mut Vec) -> Option> { +async fn try_other_providers( + skip: &str, + log: &mut Vec, +) -> Option<(Arc, goose_providers::model::ModelConfig)> { for (meta, _) in providers::providers().await { if meta.name == skip { continue; @@ -210,16 +218,21 @@ async fn try_other_providers(skip: &str, log: &mut Vec) -> Option e, Err(_) => continue, }; + let model_name = entry.metadata().default_model.clone(); + let model_config = + match crate::model_config::model_config_from_user_config(&meta.name, &model_name) { + Ok(config) => config, + Err(_) => continue, + }; let provider = match entry.create_with_default_model(vec![]).await { Ok(p) => p, Err(_) => continue, }; - let model_name = provider.get_model_config().model_name.clone(); log.push(format!(" Trying {} / {} ...", meta.name, model_name)); - match test_provider(provider.as_ref()).await { + match test_provider(provider.as_ref(), &model_config).await { Ok(()) => { log.push(format!(" ✓ {} / {} works", meta.name, model_name)); - return Some(provider); + return Some((provider, model_config)); } Err(e) => log.push(format!(" ✗ {}", describe_error(&e))), } diff --git a/crates/goose/src/execution/manager.rs b/crates/goose/src/execution/manager.rs index 68ce714da4c1..b2bb1f7bf279 100644 --- a/crates/goose/src/execution/manager.rs +++ b/crates/goose/src/execution/manager.rs @@ -238,8 +238,21 @@ impl AgentManager { if agent.provider().await.is_err() { if let Some(provider) = &*self.default_provider.read().await { + let config = crate::config::Config::global(); + let model_config = config + .get_goose_provider() + .ok() + .zip(config.get_goose_model().ok()) + .and_then(|(provider_name, model_name)| { + crate::model_config::model_config_from_user_config( + &provider_name, + &model_name, + ) + .ok() + }) + .unwrap_or_else(|| goose_providers::model::ModelConfig::new("unknown")); agent - .update_provider(Arc::clone(provider), session_id) + .update_provider(Arc::clone(provider), model_config, session_id) .await?; provider .update_mode(session_id, mode) @@ -622,10 +635,6 @@ mod tests { "failing-test-provider" } - fn get_model_config(&self) -> ModelConfig { - ModelConfig::new_or_fail("test-model") - } - async fn stream( &self, _model_config: &ModelConfig, diff --git a/crates/goose/src/model_config.rs b/crates/goose/src/model_config.rs index 779ab3e9aa2a..2e616de76dfc 100644 --- a/crates/goose/src/model_config.rs +++ b/crates/goose/src/model_config.rs @@ -1,6 +1,11 @@ use crate::config::{Config, ConfigError}; +use crate::conversation::message::Message; +use crate::providers::base::Provider; use anyhow::{anyhow, Result}; +use goose_providers::conversation::token_usage::ProviderUsage; +use goose_providers::errors::ProviderError; use goose_providers::model::ModelConfig; +use rmcp::model::Tool; use serde_json::Value; use std::collections::HashMap; @@ -59,23 +64,83 @@ fn materialize_model_config_inner( Ok(model) } -pub fn configured_fast_model_name(default_model: &str) -> String { +fn configured_fast_model_name() -> Option { Config::global() .get_param::("GOOSE_FAST_MODEL") .ok() .map(|v| v.trim().to_string()) .filter(|v| !v.is_empty()) - .unwrap_or_else(|| default_model.to_string()) } -pub fn with_configured_fast_model( - model: ModelConfig, +/// Resolve the model config to use for lightweight "fast" tasks (session +/// naming, compaction, summarization). Resolution order: +/// 1. `GOOSE_FAST_MODEL` (user override) +/// 2. the provider's declared default fast model +/// 3. the supplied `model_config` (i.e. the main model) +/// +/// The resulting config is materialized against the same provider so it picks +/// up context limits, temperature, and other provider defaults. +pub async fn get_fast_model( provider_name: &str, - default_fast_model_name: &str, + model_config: &ModelConfig, ) -> Result { - let fast_model_name = configured_fast_model_name(default_fast_model_name); - let fast_model_config = model_config_from_user_config(provider_name, fast_model_name)?; - Ok(model.with_fast_model_config(fast_model_config)) + let fast_model_name = match configured_fast_model_name() { + Some(name) => Some(name), + None => provider_default_fast_model(provider_name).await, + }; + + match fast_model_name { + Some(name) if name != model_config.model_name => { + model_config_from_user_config(provider_name, name) + } + _ => Ok(model_config.clone()), + } +} + +/// Run a completion for a lightweight "fast" task (session naming, compaction, +/// summarization) using the provider's fast model, falling back to the supplied +/// main `model_config` if the fast model errors. +pub async fn complete_fast( + provider: &dyn Provider, + model_config: &ModelConfig, + session_id: &str, + system: &str, + messages: &[Message], + tools: &[Tool], +) -> Result<(Message, ProviderUsage), ProviderError> { + let fast_model_config = get_fast_model(provider.get_name(), model_config) + .await + .map_err(|e| ProviderError::ExecutionError(e.to_string()))?; + + match provider + .complete(&fast_model_config, session_id, system, messages, tools) + .await + { + Ok(response) => Ok(response), + Err(e) if fast_model_config.model_name != model_config.model_name => { + tracing::warn!( + "Fast model {} failed with error: {}. Falling back to main model {}", + fast_model_config.model_name, + e, + model_config.model_name + ); + provider + .complete(model_config, session_id, system, messages, tools) + .await + } + Err(e) => Err(e), + } +} + +async fn provider_default_fast_model(provider_name: &str) -> Option { + if provider_name == goose_providers::openai::OPEN_AI_PROVIDER_NAME { + return crate::providers::openai_def::live_fast_model(); + } + + crate::providers::get_from_registry(provider_name) + .await + .ok() + .and_then(|entry| entry.metadata().fast_model.clone()) } fn base_model_config_from_user_config(model_name: &str) -> Result { @@ -87,7 +152,6 @@ fn base_model_config_from_user_config(model_name: &str) -> Result { max_tokens: None, toolshim: get_goose_toolshim(config)?.unwrap_or(false), toolshim_model: get_goose_toolshim_model(config)?, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -114,6 +178,14 @@ fn get_goose_toolshim(config: &Config) -> Result> { } } +/// Resolve the global toolshim setting, defaulting to false when unset. +pub fn global_toolshim() -> bool { + get_goose_toolshim(Config::global()) + .ok() + .flatten() + .unwrap_or(false) +} + fn get_goose_toolshim_model(config: &Config) -> Result> { match config.get_param::("GOOSE_TOOLSHIM_OLLAMA_MODEL") { Ok(value) if value.trim().is_empty() => Err(anyhow!( diff --git a/crates/goose/src/permission/permission_inspector.rs b/crates/goose/src/permission/permission_inspector.rs index 6f62f5609673..d7d1c8c87407 100644 --- a/crates/goose/src/permission/permission_inspector.rs +++ b/crates/goose/src/permission/permission_inspector.rs @@ -15,14 +15,20 @@ use std::sync::{Arc, RwLock}; pub struct PermissionInspector { pub permission_manager: Arc, provider: SharedProvider, + session_manager: Arc, readonly_tools: RwLock>, } impl PermissionInspector { - pub fn new(permission_manager: Arc, provider: SharedProvider) -> Self { + pub fn new( + permission_manager: Arc, + provider: SharedProvider, + session_manager: Arc, + ) -> Self { Self { permission_manager, provider, + session_manager, readonly_tools: RwLock::new(HashSet::new()), } } @@ -212,12 +218,15 @@ impl ToolInspector for PermissionInspector { // LLM-based read-only detection for deferred SmartApprove candidates if !llm_detect_candidates.is_empty() { let detected: HashSet = match self.provider.lock().await.clone() { - Some(provider) => { - detect_read_only_tools(provider, session_id, llm_detect_candidates.to_vec()) - .await - .into_iter() - .collect() - } + Some(provider) => detect_read_only_tools( + provider, + &self.session_manager, + session_id, + llm_detect_candidates.to_vec(), + ) + .await + .into_iter() + .collect(), None => Default::default(), }; @@ -288,7 +297,10 @@ mod tests { if let Some(level) = cache { pm.update_smart_approve_permission("tool", level); } - let inspector = PermissionInspector::new(pm, Arc::new(Mutex::new(None))); + let session_manager = Arc::new(crate::session::SessionManager::new( + tempfile::tempdir().unwrap().keep(), + )); + let inspector = PermissionInspector::new(pm, Arc::new(Mutex::new(None)), session_manager); if smart_approved { *inspector.readonly_tools.write().unwrap() = ["tool".to_string()].into_iter().collect(); } diff --git a/crates/goose/src/permission/permission_judge.rs b/crates/goose/src/permission/permission_judge.rs index b6dbe3c96f6f..59db44f1dc62 100644 --- a/crates/goose/src/permission/permission_judge.rs +++ b/crates/goose/src/permission/permission_judge.rs @@ -10,6 +10,28 @@ use serde::{Deserialize, Serialize}; use serde_json::Value; use std::sync::Arc; +async fn resolve_model_config( + session_manager: &crate::session::SessionManager, + session_id: &str, +) -> anyhow::Result { + if !session_id.is_empty() { + if let Ok(session) = session_manager.get_session(session_id, false).await { + if let Some(model_config) = session.model_config { + return Ok(model_config); + } + } + } + + let config = crate::config::Config::global(); + let provider_name = config + .get_goose_provider() + .map_err(|_| anyhow::anyhow!("missing provider"))?; + let model_name = config + .get_goose_model() + .map_err(|_| anyhow::anyhow!("missing model"))?; + crate::model_config::model_config_from_user_config(&provider_name, &model_name) +} + #[derive(Serialize)] struct PermissionJudgeContext { // Empty struct for now since the current template doesn't need variables @@ -121,6 +143,7 @@ fn extract_read_only_tools(response: &Message) -> Option> { /// Executes the read-only tools detection and returns the list of tools with read-only operations. pub async fn detect_read_only_tools( provider: Arc, + session_manager: &crate::session::SessionManager, session_id: &str, tool_requests: Vec<&ToolRequest>, ) -> Vec { @@ -134,7 +157,13 @@ pub async fn detect_read_only_tools( let system_prompt = render_template("permission_judge.md", &context) .unwrap_or_else(|_| "You are a good analyst and can detect operations whether they have read-only operations.".to_string()); - let model_config = provider.get_model_config(); + let model_config = match resolve_model_config(session_manager, session_id).await { + Ok(config) => config, + Err(e) => { + tracing::warn!("Could not resolve model config for permission judge: {e}"); + return vec![]; + } + }; let res = provider .complete( &model_config, diff --git a/crates/goose/src/providers/amp_acp.rs b/crates/goose/src/providers/amp_acp.rs index beae52173ce7..f32c662c3321 100644 --- a/crates/goose/src/providers/amp_acp.rs +++ b/crates/goose/src/providers/amp_acp.rs @@ -11,7 +11,6 @@ use crate::config::{Config, GooseMode}; use crate::providers::base::{ current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, }; -use goose_providers::model::ModelConfig; pub(crate) const AMP_ACP_PROVIDER_NAME: &str = "amp-acp"; const AMP_ACP_DOC_URL: &str = "https://ampcode.com"; @@ -45,15 +44,13 @@ impl ProviderDef for AmpAcpProvider { type Provider = AcpProvider; fn from_env( - model: ModelConfig, extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Self::from_env_with_working_dir(model, extensions, current_working_dir(), tls_config) + Self::from_env_with_working_dir(extensions, current_working_dir(), tls_config) } fn from_env_with_working_dir( - model: ModelConfig, extensions: Vec, working_dir: PathBuf, _tls_config: Option, @@ -80,12 +77,14 @@ impl ProviderDef for AmpAcpProvider { work_dir: working_dir, mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), + session_config_options: vec![], + model_config_option_id: None, mode_mapping, notification_callback: None, }; let metadata = Self::metadata(); - AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await + AcpProvider::connect(metadata.name, goose_mode, provider_config).await }) } } diff --git a/crates/goose/src/providers/anthropic.rs b/crates/goose/src/providers/anthropic.rs index 15d2d53af349..80379a444ceb 100644 --- a/crates/goose/src/providers/anthropic.rs +++ b/crates/goose/src/providers/anthropic.rs @@ -14,8 +14,8 @@ use tokio_util::io::StreamReader; use super::api_client::{ApiClient, AuthMethod}; use super::base::{ConfigKey, MessageStream, ModelInfo, Provider, ProviderDef, ProviderMetadata}; use super::formats::anthropic::{ - create_request_with_options_for_provider, response_to_streaming_message, thinking_type, - AnthropicFormatOptions, ThinkingType, ANTHROPIC_PROVIDER_NAME, + create_request_with_options_for_provider, response_to_streaming_message, + AnthropicFormatOptions, ANTHROPIC_PROVIDER_NAME, }; use super::openai_compatible::handle_status; use super::openai_compatible::map_http_error_to_provider_error; @@ -27,7 +27,6 @@ use goose_providers::model::ModelConfig; use rmcp::model::Tool; pub const ANTHROPIC_DEFAULT_MODEL: &str = "claude-sonnet-4-5"; -const ANTHROPIC_DEFAULT_FAST_MODEL: &str = "claude-haiku-4-5"; const ANTHROPIC_KNOWN_MODELS: &[&str] = &[ // Claude 4.6 models "claude-opus-4-6", @@ -48,12 +47,12 @@ const ANTHROPIC_KNOWN_MODELS: &[&str] = &[ const ANTHROPIC_DOC_URL: &str = "https://docs.anthropic.com/en/docs/about-claude/models"; const ANTHROPIC_API_VERSION: &str = "2023-06-01"; +const ANTHROPIC_DEFAULT_FAST_MODEL: &str = "claude-haiku-4-5"; #[derive(serde::Serialize)] pub struct AnthropicProvider { #[serde(skip)] api_client: ApiClient, - model: ModelConfig, supports_streaming: bool, name: String, custom_models: Option>, @@ -65,15 +64,8 @@ pub struct AnthropicProvider { impl AnthropicProvider { pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { - let model = crate::model_config::with_configured_fast_model( - model, - ANTHROPIC_PROVIDER_NAME, - ANTHROPIC_DEFAULT_FAST_MODEL, - )?; - let config = crate::config::Config::global(); let api_key: String = config.get_secret("ANTHROPIC_API_KEY")?; let host: String = config @@ -90,7 +82,6 @@ impl AnthropicProvider { Ok(Self { api_client, - model, supports_streaming: true, name: ANTHROPIC_PROVIDER_NAME.to_string(), custom_models: None, @@ -101,7 +92,6 @@ impl AnthropicProvider { } pub fn from_custom_config( - model: ModelConfig, config: DeclarativeProviderConfig, tls_config: Option, ) -> Result { @@ -159,15 +149,8 @@ impl AnthropicProvider { )); } - let model = if let Some(ref fast_model_name) = config.fast_model { - crate::model_config::with_configured_fast_model(model, &config.name, fast_model_name)? - } else { - model - }; - Ok(Self { api_client, - model, supports_streaming, name: config.name.clone(), custom_models, @@ -184,19 +167,6 @@ impl AnthropicProvider { } } - fn get_conditional_headers(&self) -> Vec<(&str, &str)> { - let mut headers = Vec::new(); - - if self.model.model_name.starts_with("claude-3-7-sonnet-") { - if thinking_type(&self.model) == ThinkingType::Enabled { - headers.push(("anthropic-beta", "output-128k-2025-02-19")); - } - headers.push(("anthropic-beta", "token-efficient-tools-2025-02-19")); - } - - headers - } - async fn fetch_models_from_api(&self) -> Result, ProviderError> { let response = self.api_client.request(None, "v1/models").api_get().await?; @@ -265,6 +235,7 @@ impl ProviderDescriptor for AnthropicProvider { "Click 'Create Key'", "Copy the key and paste it above", ]) + .with_fast_model(ANTHROPIC_DEFAULT_FAST_MODEL) } } @@ -272,11 +243,10 @@ impl ProviderDef for AnthropicProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -290,10 +260,6 @@ impl Provider for AnthropicProvider { self.skip_canonical_filtering } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { if let Some(custom_models) = &self.custom_models { if self.dynamic_models == Some(false) { @@ -337,15 +303,11 @@ impl Provider for AnthropicProvider { .unwrap() .insert("stream".to_string(), Value::Bool(true)); - let conditional_headers = self.get_conditional_headers(); let mut log = start_log(model_config, &payload)?; let response = self .with_retry(|| async { - let mut request = self.api_client.request(Some(session_id), "v1/messages"); - for (key, value) in &conditional_headers { - request = request.header(key, value)?; - } + let request = self.api_client.request(Some(session_id), "v1/messages"); let resp = request.response_post(&payload).await?; handle_status(resp).await }) @@ -393,7 +355,6 @@ mod tests { .unwrap(); AnthropicProvider { api_client, - model: ModelConfig::new_or_fail("claude-test"), supports_streaming: true, name: "custom_anthropic".to_string(), custom_models, @@ -452,13 +413,9 @@ mod tests { #[test] fn from_custom_config_rejects_static_only_without_models() { let config = base_declarative_config(vec![], Some(false)); - let err = AnthropicProvider::from_custom_config( - ModelConfig::new_or_fail("claude-test"), - config, - None, - ) - .err() - .expect("expected construction error for dynamic_models: false with empty models"); + let err = AnthropicProvider::from_custom_config(config, None) + .err() + .expect("expected construction error for dynamic_models: false with empty models"); let msg = err.to_string(); assert!( msg.contains("dynamic_models: false"), diff --git a/crates/goose/src/providers/avian.rs b/crates/goose/src/providers/avian.rs index 17c8e777933c..05105a1c24bf 100644 --- a/crates/goose/src/providers/avian.rs +++ b/crates/goose/src/providers/avian.rs @@ -3,7 +3,6 @@ use super::base::{ConfigKey, ProviderDef, ProviderMetadata}; use super::openai_compatible::OpenAiCompatibleProvider; use anyhow::Result; use futures::future::BoxFuture; -use goose_providers::model::ModelConfig; const AVIAN_PROVIDER_NAME: &str = "avian"; pub const AVIAN_API_HOST: &str = "https://api.avian.io/v1"; @@ -39,7 +38,6 @@ impl ProviderDef for AvianProvider { type Provider = OpenAiCompatibleProvider; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { @@ -56,7 +54,6 @@ impl ProviderDef for AvianProvider { Ok(OpenAiCompatibleProvider::new( AVIAN_PROVIDER_NAME.to_string(), api_client, - model, String::new(), )) }) diff --git a/crates/goose/src/providers/azure.rs b/crates/goose/src/providers/azure.rs index c217b4dbf166..f0f10ec9466f 100644 --- a/crates/goose/src/providers/azure.rs +++ b/crates/goose/src/providers/azure.rs @@ -6,7 +6,6 @@ use super::azureauth::{AuthError, AzureAuth}; use super::base::{ConfigKey, ProviderDef, ProviderMetadata}; use super::openai_compatible::OpenAiCompatibleProvider; use futures::future::BoxFuture; -use goose_providers::model::ModelConfig; const AZURE_PROVIDER_NAME: &str = "azure_openai"; pub const AZURE_DEFAULT_MODEL: &str = "gpt-4o"; @@ -74,7 +73,6 @@ impl ProviderDef for AzureProvider { type Provider = OpenAiCompatibleProvider; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { @@ -120,7 +118,6 @@ impl ProviderDef for AzureProvider { Ok(OpenAiCompatibleProvider::new( AZURE_PROVIDER_NAME.to_string(), api_client, - model, format!("deployments/{}/", deployment_name), )) }) diff --git a/crates/goose/src/providers/base.rs b/crates/goose/src/providers/base.rs index c42c97e52c2a..4f3c096c0521 100644 --- a/crates/goose/src/providers/base.rs +++ b/crates/goose/src/providers/base.rs @@ -9,7 +9,6 @@ use serde::{Deserialize, Serialize}; pub const DEFAULT_PROVIDER_TIMEOUT_SECS: u64 = 600; use crate::config::ExtensionConfig; -use goose_providers::model::ModelConfig; use utoipa::ToSchema; use std::path::PathBuf; @@ -32,7 +31,6 @@ pub trait ProviderDef: ProviderDescriptor + Send + Sync { type Provider: Provider + 'static; fn from_env( - model: ModelConfig, extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> @@ -40,7 +38,6 @@ pub trait ProviderDef: ProviderDescriptor + Send + Sync { Self: Sized; fn from_env_with_working_dir( - model: ModelConfig, extensions: Vec, _working_dir: PathBuf, tls_config: Option, @@ -48,6 +45,6 @@ pub trait ProviderDef: ProviderDescriptor + Send + Sync { where Self: Sized, { - Self::from_env(model, extensions, tls_config) + Self::from_env(extensions, tls_config) } } diff --git a/crates/goose/src/providers/bedrock.rs b/crates/goose/src/providers/bedrock.rs index c06347e07f89..b23bb89b90d3 100644 --- a/crates/goose/src/providers/bedrock.rs +++ b/crates/goose/src/providers/bedrock.rs @@ -55,7 +55,6 @@ pub const BEDROCK_DEFAULT_MAX_RETRY_INTERVAL_MS: u64 = 120_000; pub struct BedrockProvider { #[serde(skip)] client: Client, - model: ModelConfig, #[serde(skip)] retry_config: RetryConfig, #[serde(skip)] @@ -80,7 +79,6 @@ struct ConverseRequestParts { impl BedrockProvider { pub async fn from_env( - model: ModelConfig, _tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -173,7 +171,6 @@ impl BedrockProvider { Ok(Self { client, - model, retry_config, name: BEDROCK_PROVIDER_NAME.to_string(), region: resolved_region, @@ -224,13 +221,13 @@ impl BedrockProvider { ) } - fn should_enable_caching(&self) -> bool { + fn should_enable_caching(&self, model: &ModelConfig) -> bool { let config = crate::config::Config::global(); let enabled = config .get_param::("BEDROCK_ENABLE_CACHING") .unwrap_or(false); - enabled && self.model.model_name.contains("anthropic.claude") + enabled && model.model_name.contains("anthropic.claude") } async fn post_mantle_streaming( @@ -285,11 +282,12 @@ impl BedrockProvider { /// the tool configuration. fn build_request_parts( &self, + model: &ModelConfig, system: &str, messages: &[Message], tools: &[Tool], ) -> Result { - let enable_caching = self.should_enable_caching(); + let enable_caching = self.should_enable_caching(model); let system_blocks = if enable_caching { vec![ @@ -332,26 +330,25 @@ impl BedrockProvider { system_blocks, messages: bedrock_messages, tool_config, - thinking_fields: bedrock_anthropic_thinking_fields(&self.model), + thinking_fields: bedrock_anthropic_thinking_fields(model), }) } async fn converse( &self, + model: &ModelConfig, session_id: Option<&str>, system: &str, messages: &[Message], tools: &[Tool], ) -> Result<(bedrock::Message, Option), ProviderError> { - let model_name = &self.model.model_name; - - let parts = self.build_request_parts(system, messages, tools)?; + let parts = self.build_request_parts(model, system, messages, tools)?; let mut request = self .client .converse() .set_system(Some(parts.system_blocks)) - .model_id(model_name.to_string()) + .model_id(&model.model_name) .set_messages(Some(parts.messages)); if let Some(fields) = parts.thinking_fields { @@ -431,6 +428,7 @@ impl BedrockProvider { /// receiver so [`Provider::stream`] can forward deltas incrementally. async fn converse_stream( &self, + model: &ModelConfig, session_id: Option<&str>, system: &str, messages: &[Message], @@ -439,15 +437,13 @@ impl BedrockProvider { aws_sdk_bedrockruntime::operation::converse_stream::ConverseStreamOutput, ProviderError, > { - let model_name = &self.model.model_name; - - let parts = self.build_request_parts(system, messages, tools)?; + let parts = self.build_request_parts(model, system, messages, tools)?; let mut request = self .client .converse_stream() .set_system(Some(parts.system_blocks)) - .model_id(model_name.to_string()) + .model_id(&model.model_name) .set_messages(Some(parts.messages)); if let Some(fields) = parts.thinking_fields { @@ -506,6 +502,7 @@ impl BedrockProvider { /// escape hatch. async fn stream_via_converse( &self, + model: &ModelConfig, session_id: Option<&str>, system: &str, messages: &[Message], @@ -513,7 +510,7 @@ impl BedrockProvider { model_name: &str, ) -> Result { let (bedrock_message, bedrock_usage) = self - .with_retry(|| self.converse(session_id, system, messages, tools)) + .with_retry(|| self.converse(model, session_id, system, messages, tools)) .await?; let usage = bedrock_usage @@ -529,7 +526,7 @@ impl BedrockProvider { "messages": messages, "tools": tools }); - let mut log = start_log(&self.model, &debug_payload)?; + let mut log = start_log(model, &debug_payload)?; log.write( &serde_json::to_value(&message).unwrap_or_default(), Some(&usage), @@ -713,11 +710,10 @@ impl ProviderDef for BedrockProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -731,10 +727,6 @@ impl Provider for BedrockProvider { self.retry_config.clone() } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { Ok(BEDROCK_KNOWN_MODELS.iter().map(|s| s.to_string()).collect()) } @@ -798,7 +790,14 @@ impl Provider for BedrockProvider { // Escape hatch: restore the previous blocking-Converse behaviour. if self.streaming_disabled() { return self - .stream_via_converse(session_id_opt, system, messages, tools, &model_name) + .stream_via_converse( + model_config, + session_id_opt, + system, + messages, + tools, + &model_name, + ) .await; } @@ -806,7 +805,9 @@ impl Provider for BedrockProvider { // setup only — mid-stream errors are surfaced, not retried (matching // the Anthropic provider's behaviour). let response = self - .with_retry(|| self.converse_stream(session_id_opt, system, messages, tools)) + .with_retry(|| { + self.converse_stream(model_config, session_id_opt, system, messages, tools) + }) .await?; // Debug trace with input context; the streamed text is written once @@ -816,7 +817,7 @@ impl Provider for BedrockProvider { "messages": messages, "tools": tools }); - let mut log = start_log(&self.model, &debug_payload)?; + let mut log = start_log(model_config, &debug_payload)?; let mut event_stream = response.stream; @@ -909,33 +910,34 @@ mod tests { use goose_providers::base::ProviderDescriptor as _; use serial_test::serial; - fn create_mock_provider(model_name: &str) -> BedrockProvider { + fn create_mock_provider_and_model(model_name: &str) -> (BedrockProvider, ModelConfig) { let sdk_config = aws_config::SdkConfig::builder() .behavior_version(aws_config::BehaviorVersion::latest()) .region(aws_config::Region::new("us-east-1")) .build(); let client = Client::new(&sdk_config); - BedrockProvider { - client, - model: ModelConfig { + ( + BedrockProvider { + client, + retry_config: RetryConfig::default(), + name: "aws_bedrock".to_string(), + region: None, + bearer_token: None, + http_client: reqwest::Client::new(), + mantle_base_url: None, + }, + ModelConfig { model_name: model_name.to_string(), context_limit: None, temperature: None, max_tokens: None, toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }, - retry_config: RetryConfig::default(), - name: "aws_bedrock".to_string(), - region: None, - bearer_token: None, - http_client: reqwest::Client::new(), - mantle_base_url: None, - } + ) } #[test] @@ -1000,24 +1002,11 @@ mod tests { ); } - #[test] - #[serial] - fn test_caching_disabled_by_default() { - // Ensure clean environment - std::env::remove_var("BEDROCK_ENABLE_CACHING"); - - let provider = create_mock_provider("us.anthropic.claude-sonnet-4-5-20250929-v1:0"); - assert!( - !provider.should_enable_caching(), - "Caching should be disabled by default" - ); - } - #[test] fn test_caching_disabled_for_non_claude_models() { - let provider = create_mock_provider("amazon.titan-text-express-v1"); + let (provider, model) = create_mock_provider_and_model("amazon.titan-text-express-v1"); assert!( - !provider.should_enable_caching(), + !provider.should_enable_caching(&model), "Caching should be disabled for non-Claude models" ); } @@ -1027,9 +1016,10 @@ mod tests { fn test_caching_enabled_for_claude_model() { std::env::set_var("BEDROCK_ENABLE_CACHING", "true"); - let provider = create_mock_provider("us.anthropic.claude-sonnet-4-5-20250929-v1:0"); + let (provider, model) = + create_mock_provider_and_model("us.anthropic.claude-sonnet-4-5-20250929-v1:0"); assert!( - provider.should_enable_caching(), + provider.should_enable_caching(&model), "Caching should be enabled for Claude models when BEDROCK_ENABLE_CACHING=true" ); @@ -1038,7 +1028,7 @@ mod tests { #[tokio::test] async fn test_post_mantle_streaming_missing_region() { - let provider = create_mock_provider("openai.gpt-5.5"); + let (provider, _) = create_mock_provider_and_model("openai.gpt-5.5"); let payload = serde_json::json!({"model": "openai.gpt-5.5"}); let result = provider.post_mantle_streaming(None, &payload).await; assert!(result.is_err()); @@ -1055,7 +1045,7 @@ mod tests { #[tokio::test] async fn test_post_mantle_streaming_missing_bearer_token() { - let mut provider = create_mock_provider("openai.gpt-5.5"); + let (mut provider, _) = create_mock_provider_and_model("openai.gpt-5.5"); provider.region = Some("us-east-1".to_string()); let payload = serde_json::json!({"model": "openai.gpt-5.5"}); @@ -1099,9 +1089,9 @@ mod tests { .region(aws_config::Region::new("us-east-1")) .build(); + let model = ModelConfig::new("openai.gpt-5.5"); let provider = BedrockProvider { client: Client::new(&sdk_config), - model: ModelConfig::new("openai.gpt-5.5").unwrap(), retry_config: RetryConfig::default(), name: "aws_bedrock".to_string(), region: Some("us-east-1".to_string()), @@ -1112,7 +1102,7 @@ mod tests { let messages = vec![crate::conversation::message::Message::user().with_text("hi")]; let mut stream = provider - .stream(&provider.model.clone(), "", "", &messages, &[]) + .stream(&model.clone(), "", "", &messages, &[]) .await .unwrap(); diff --git a/crates/goose/src/providers/chatgpt_codex.rs b/crates/goose/src/providers/chatgpt_codex.rs index ff6213d86b18..cb74d8837a7c 100644 --- a/crates/goose/src/providers/chatgpt_codex.rs +++ b/crates/goose/src/providers/chatgpt_codex.rs @@ -879,7 +879,6 @@ impl AuthProvider for ChatGptCodexAuthProvider { pub struct ChatGptCodexProvider { #[serde(skip)] auth_provider: Arc, - model: ModelConfig, #[serde(skip)] name: String, } @@ -891,7 +890,6 @@ impl ChatGptCodexProvider { } pub async fn from_env( - model: ModelConfig, _tls_config: Option, ) -> Result { let auth_provider = Arc::new(ChatGptCodexAuthProvider::new( @@ -900,7 +898,6 @@ impl ChatGptCodexProvider { Ok(Self { auth_provider, - model, name: CHATGPT_CODEX_PROVIDER_NAME.to_string(), }) } @@ -975,11 +972,10 @@ impl ProviderDef for ChatGptCodexProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -989,10 +985,6 @@ impl Provider for ChatGptCodexProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn stream( &self, model_config: &ModelConfig, @@ -1214,7 +1206,7 @@ mod tests { fn test_create_codex_request_reasoning_effort_from_unified_thinking() { let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), json!("max")); - let mut config = ModelConfig::new("gpt-5.3-codex").unwrap(); + let mut config = ModelConfig::new("gpt-5.3-codex"); config.request_params = Some(params); let payload = create_codex_request(&config, "sys", &[], &[]).unwrap(); @@ -1226,7 +1218,7 @@ mod tests { fn test_create_codex_request_caps_unified_thinking_to_supported_level() { let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), json!("max")); - let mut config = ModelConfig::new("unknown-model").unwrap(); + let mut config = ModelConfig::new("unknown-model"); config.request_params = Some(params); let payload = create_codex_request(&config, "sys", &[], &[]).unwrap(); @@ -1238,7 +1230,7 @@ mod tests { fn test_create_codex_request_off_omits_reasoning_for_codex_models() { let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), json!("off")); - let mut config = ModelConfig::new("gpt-5.2-codex").unwrap(); + let mut config = ModelConfig::new("gpt-5.2-codex"); config.request_params = Some(params); let payload = create_codex_request(&config, "sys", &[], &[]).unwrap(); @@ -1249,9 +1241,7 @@ mod tests { // ChatGPT Codex does not support temperature and will return an error #[test] fn test_create_codex_request_omits_temperature() { - let config = ModelConfig::new("gpt-5.5") - .unwrap() - .with_temperature(Some(0.2)); + let config = ModelConfig::new("gpt-5.5").with_temperature(Some(0.2)); let payload = create_codex_request(&config, "sys", &[], &[]).unwrap(); assert!(payload.get("temperature").is_none()); @@ -1417,7 +1407,7 @@ mod tests { #[test] fn test_gpt53_preamble_injected() { - let model = ModelConfig::new("gpt-5.3-codex").unwrap(); + let model = ModelConfig::new("gpt-5.3-codex"); let payload = create_codex_request(&model, "system prompt", &[], &[]).unwrap(); let instructions = payload["instructions"].as_str().unwrap(); assert!(instructions.contains(GPT_53_CODEX_TOOL_PREAMBLE)); @@ -1426,7 +1416,7 @@ mod tests { #[test] fn test_other_models_no_preamble() { - let model = ModelConfig::new("gpt-5.4").unwrap(); + let model = ModelConfig::new("gpt-5.4"); let payload = create_codex_request(&model, "system prompt", &[], &[]).unwrap(); let instructions = payload["instructions"].as_str().unwrap(); assert_eq!(instructions, "system prompt"); diff --git a/crates/goose/src/providers/claude_acp.rs b/crates/goose/src/providers/claude_acp.rs index f3ab9d0e229e..46efe7d5386c 100644 --- a/crates/goose/src/providers/claude_acp.rs +++ b/crates/goose/src/providers/claude_acp.rs @@ -11,7 +11,6 @@ use crate::config::{Config, GooseMode}; use crate::providers::base::{ current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, }; -use goose_providers::model::ModelConfig; pub(crate) const CLAUDE_ACP_PROVIDER_NAME: &str = "claude-acp"; const CLAUDE_ACP_DOC_URL: &str = "https://github.com/agentclientprotocol/claude-agent-acp"; @@ -43,15 +42,13 @@ impl ProviderDef for ClaudeAcpProvider { type Provider = AcpProvider; fn from_env( - model: ModelConfig, extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Self::from_env_with_working_dir(model, extensions, current_working_dir(), tls_config) + Self::from_env_with_working_dir(extensions, current_working_dir(), tls_config) } fn from_env_with_working_dir( - model: ModelConfig, extensions: Vec, working_dir: PathBuf, _tls_config: Option, @@ -84,12 +81,14 @@ impl ProviderDef for ClaudeAcpProvider { work_dir: working_dir, mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), + session_config_options: vec![], + model_config_option_id: None, mode_mapping, notification_callback: None, }; let metadata = Self::metadata(); - AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await + AcpProvider::connect(metadata.name, goose_mode, provider_config).await }) } } diff --git a/crates/goose/src/providers/claude_code.rs b/crates/goose/src/providers/claude_code.rs index 55377997efb1..4adbc867f7f6 100644 --- a/crates/goose/src/providers/claude_code.rs +++ b/crates/goose/src/providers/claude_code.rs @@ -259,7 +259,6 @@ impl Drop for CliProcess { #[derive(Debug, serde::Serialize)] pub struct ClaudeCodeProvider { command: PathBuf, - model: ModelConfig, #[serde(skip)] name: String, /// Temp file holding MCP config JSON (auto-deleted on drop). @@ -364,7 +363,11 @@ impl ClaudeCodeProvider { } } - async fn spawn_process(&self, filtered_system: &str) -> Result { + async fn spawn_process( + &self, + model: &ModelConfig, + filtered_system: &str, + ) -> Result { let mut cmd = self.build_stream_json_command(); if let Some(f) = &self.mcp_config_file { @@ -376,7 +379,7 @@ impl ClaudeCodeProvider { .arg("--system-prompt") .arg(filtered_system) .arg("--model") - .arg(&self.model.model_name); + .arg(&model.model_name); let control_protocol_enabled = Self::apply_permission_flags(&mut cmd)?; @@ -411,7 +414,7 @@ impl ClaudeCodeProvider { stdin: Box::new(stdin), reader: BufReader::new(Box::new(stdout)), stderr_handle, - current_model: self.model.model_name.clone(), + current_model: model.model_name.clone(), log_model_update: false, next_request_id: 0, needs_drain: false, @@ -428,12 +431,13 @@ impl ClaudeCodeProvider { async fn get_or_init_process( &self, + model_config: &ModelConfig, filtered_system: &str, ) -> Result<&Arc>, ProviderError> { self.cli_process .get_or_try_init(|| async { Ok(Arc::new(tokio::sync::Mutex::new( - self.spawn_process(filtered_system).await?, + self.spawn_process(model_config, filtered_system).await?, ))) }) .await @@ -607,7 +611,6 @@ impl ProviderDef for ClaudeCodeProvider { type Provider = Self; fn from_env( - model: ModelConfig, extensions: Vec, _tls_config: Option, ) -> BoxFuture<'static, Result> { @@ -627,7 +630,6 @@ impl ProviderDef for ClaudeCodeProvider { Ok(Self { command: resolved_command, - model, name: CLAUDE_CODE_PROVIDER_NAME.to_string(), mcp_config_file, cli_process: tokio::sync::OnceCell::new(), @@ -648,10 +650,6 @@ impl Provider for ClaudeCodeProvider { true } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { // Uses a separate short-lived process because --system-prompt is a CLI-only // flag with no NDJSON equivalent. The persistent process needs it at spawn, @@ -730,7 +728,10 @@ impl Provider for ClaudeCodeProvider { } let filtered_system = filter_extensions_from_system_prompt(system); - let process_arc = Arc::clone(self.get_or_init_process(&filtered_system).await?); + let process_arc = Arc::clone( + self.get_or_init_process(model_config, &filtered_system) + .await?, + ); // Prepare the payload outside the lock — these don't need the process. let blocks = self.last_user_content_blocks(messages); @@ -1275,9 +1276,6 @@ mod tests { fn make_provider() -> ClaudeCodeProvider { ClaudeCodeProvider { command: PathBuf::from("claude"), - model: ModelConfig::new(CLAUDE_CODE_DEFAULT_MODEL) - .unwrap() - .with_canonical_limits(CLAUDE_CODE_PROVIDER_NAME), name: "claude-code".to_string(), mcp_config_file: None, cli_process: tokio::sync::OnceCell::new(), @@ -1316,8 +1314,10 @@ mod tests { provider.cli_process.set(process_arc).unwrap(); let messages = vec![Message::user().with_text("test")]; + let model = ModelConfig::new(CLAUDE_CODE_DEFAULT_MODEL) + .with_canonical_limits(CLAUDE_CODE_PROVIDER_NAME); let stream = provider - .stream(&provider.model, "test-session", "", &messages, &[]) + .stream(&model, "test-session", "", &messages, &[]) .await .unwrap(); (provider, stream, stdin_reader) @@ -1524,8 +1524,10 @@ mod tests { .insert("stale_1".to_string(), tx); let messages = vec![Message::user().with_text("test")]; + let model = ModelConfig::new(CLAUDE_CODE_DEFAULT_MODEL) + .with_canonical_limits(CLAUDE_CODE_PROVIDER_NAME); let mut stream = provider - .stream(&provider.model, "test-session", "", &messages, &[]) + .stream(&model, "test-session", "", &messages, &[]) .await .unwrap(); diff --git a/crates/goose/src/providers/codex.rs b/crates/goose/src/providers/codex.rs index fad5ce79905d..c3722640ca8a 100644 --- a/crates/goose/src/providers/codex.rs +++ b/crates/goose/src/providers/codex.rs @@ -46,11 +46,8 @@ pub const CODEX_REASONING_LEVELS: &[&str] = &["none", "low", "medium", "high", " #[derive(Debug, serde::Serialize)] pub struct CodexProvider { command: PathBuf, - model: ModelConfig, #[serde(skip)] name: String, - /// Reasoning effort level (none, low, medium, high, xhigh) - reasoning_effort: Option, /// Whether to skip git repo check skip_git_check: bool, /// CLI config overrides for MCP servers @@ -126,6 +123,7 @@ impl CodexProvider { /// Execute codex CLI command async fn execute_command( &self, + model: &ModelConfig, system: &str, messages: &[Message], _tools: &[Tool], @@ -137,10 +135,12 @@ impl CodexProvider { let (prompt, temp_files) = prepare_input(system, messages, &image_dir)?; if std::env::var("GOOSE_CODEX_DEBUG").is_ok() { + let reasoning_effort = + Self::map_thinking_effort(&model.model_name, model.thinking_effort()); println!("=== CODEX PROVIDER DEBUG ==="); println!("Command: {:?}", self.command); - println!("Model: {}", self.model.model_name); - println!("Reasoning effort: {:?}", self.reasoning_effort); + println!("Model: {}", model.model_name); + println!("Reasoning effort: {:?}", reasoning_effort); println!("Skip git check: {}", self.skip_git_check); println!("Prompt length: {} chars", prompt.len()); println!("Prompt: {}", prompt); @@ -163,11 +163,13 @@ impl CodexProvider { // Only pass model parameter if it's in the known models list // This allows users to set GOOSE_PROVIDER=codex without needing to specify a model - if CODEX_KNOWN_MODELS.contains(&self.model.model_name.as_str()) { - cmd.arg("-m").arg(&self.model.model_name); + if CODEX_KNOWN_MODELS.contains(&model.model_name.as_str()) { + cmd.arg("-m").arg(&model.model_name); } - if let Some(reasoning_effort) = &self.reasoning_effort { + if let Some(reasoning_effort) = + Self::map_thinking_effort(&model.model_name, model.thinking_effort()) + { cmd.arg("-c") .arg(format!("model_reasoning_effort=\"{}\"", reasoning_effort)); } @@ -642,7 +644,6 @@ impl ProviderDef for CodexProvider { type Provider = Self; fn from_env( - model: ModelConfig, extensions: Vec, _tls_config: Option, ) -> BoxFuture<'static, Result> { @@ -651,9 +652,6 @@ impl ProviderDef for CodexProvider { let command: String = config.get_codex_command().unwrap_or_default().into(); let resolved_command = SearchPaths::builder().with_npm().resolve(command)?; - let reasoning_effort = - Self::map_thinking_effort(&model.model_name, model.thinking_effort()); - // Get skip_git_check from config, default to false let skip_git_check = config .get_codex_skip_git_check() @@ -667,9 +665,7 @@ impl ProviderDef for CodexProvider { Ok(Self { command: resolved_command, - model, name: CODEX_PROVIDER_NAME.to_string(), - reasoning_effort, skip_git_check, mcp_config_overrides: codex_mcp_config_overrides(&resolved), mode_by_session: tokio::sync::RwLock::new(HashMap::new()), @@ -684,10 +680,6 @@ impl Provider for CodexProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn stream( &self, model_config: &ModelConfig, @@ -712,7 +704,7 @@ impl Provider for CodexProvider { map.get(session_id).copied().unwrap_or_default() }; let lines = self - .execute_command(system, messages, tools, goose_mode) + .execute_command(model_config, system, messages, tools, goose_mode) .await?; let (message, usage) = self.parse_response(&lines)?; @@ -721,7 +713,7 @@ impl Provider for CodexProvider { let payload = json!({ "command": self.command, "model": model_config.model_name, - "reasoning_effort": self.reasoning_effort, + "reasoning_effort": Self::map_thinking_effort(&model_config.model_name, model_config.thinking_effort()), "system_length": system.len(), "messages_count": messages.len() }); @@ -940,9 +932,7 @@ mod tests { fn test_parse_response_plain_text() { let provider = CodexProvider { command: PathBuf::from("codex"), - model: ModelConfig::new("gpt-5.2-codex").unwrap(), name: "codex".to_string(), - reasoning_effort: Some("high".to_string()), skip_git_check: false, mcp_config_overrides: Vec::new(), mode_by_session: tokio::sync::RwLock::new(HashMap::new()), @@ -961,9 +951,7 @@ mod tests { fn test_parse_response_json_events() { let provider = CodexProvider { command: PathBuf::from("codex"), - model: ModelConfig::new("gpt-5.2-codex").unwrap(), name: "codex".to_string(), - reasoning_effort: Some("high".to_string()), skip_git_check: false, mcp_config_overrides: Vec::new(), mode_by_session: tokio::sync::RwLock::new(HashMap::new()), @@ -996,9 +984,7 @@ mod tests { fn test_parse_response_empty() { let provider = CodexProvider { command: PathBuf::from("codex"), - model: ModelConfig::new("gpt-5.2-codex").unwrap(), name: "codex".to_string(), - reasoning_effort: Some("high".to_string()), skip_git_check: false, mcp_config_overrides: Vec::new(), mode_by_session: tokio::sync::RwLock::new(HashMap::new()), @@ -1045,9 +1031,7 @@ mod tests { fn test_parse_response_item_completed() { let provider = CodexProvider { command: PathBuf::from("codex"), - model: ModelConfig::new("gpt-5.2-codex").unwrap(), name: "codex".to_string(), - reasoning_effort: Some("high".to_string()), skip_git_check: false, mcp_config_overrides: Vec::new(), mode_by_session: tokio::sync::RwLock::new(HashMap::new()), @@ -1071,9 +1055,7 @@ mod tests { fn test_parse_response_turn_completed_usage() { let provider = CodexProvider { command: PathBuf::from("codex"), - model: ModelConfig::new("gpt-5.2-codex").unwrap(), name: "codex".to_string(), - reasoning_effort: Some("high".to_string()), skip_git_check: false, mcp_config_overrides: Vec::new(), mode_by_session: tokio::sync::RwLock::new(HashMap::new()), @@ -1145,9 +1127,7 @@ mod tests { fn test_parse_response_error_event(lines: &[&str], expected: ProviderError) { let provider = CodexProvider { command: PathBuf::from("codex"), - model: ModelConfig::new("gpt-5.2-codex").unwrap(), name: "codex".to_string(), - reasoning_effort: Some("high".to_string()), skip_git_check: false, mcp_config_overrides: Vec::new(), mode_by_session: tokio::sync::RwLock::new(HashMap::new()), @@ -1162,9 +1142,7 @@ mod tests { fn test_parse_response_skips_reasoning() { let provider = CodexProvider { command: PathBuf::from("codex"), - model: ModelConfig::new("gpt-5.2-codex").unwrap(), name: "codex".to_string(), - reasoning_effort: Some("high".to_string()), skip_git_check: false, mcp_config_overrides: Vec::new(), mode_by_session: tokio::sync::RwLock::new(HashMap::new()), @@ -1289,9 +1267,7 @@ mod tests { fn test_parse_response_multiple_agent_messages() { let provider = CodexProvider { command: PathBuf::from("codex"), - model: ModelConfig::new("gpt-5.2-codex").unwrap(), name: "codex".to_string(), - reasoning_effort: Some("high".to_string()), skip_git_check: false, mcp_config_overrides: Vec::new(), mode_by_session: tokio::sync::RwLock::new(HashMap::new()), diff --git a/crates/goose/src/providers/codex_acp.rs b/crates/goose/src/providers/codex_acp.rs index 7a3698fb63d7..aeec80cd51c3 100644 --- a/crates/goose/src/providers/codex_acp.rs +++ b/crates/goose/src/providers/codex_acp.rs @@ -11,7 +11,6 @@ use crate::config::{Config, GooseMode}; use crate::providers::base::{ current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, }; -use goose_providers::model::ModelConfig; pub(crate) const CODEX_ACP_PROVIDER_NAME: &str = "codex-acp"; const CODEX_ACP_DOC_URL: &str = "https://github.com/zed-industries/codex-acp"; @@ -42,15 +41,13 @@ impl ProviderDef for CodexAcpProvider { type Provider = AcpProvider; fn from_env( - model: ModelConfig, extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Self::from_env_with_working_dir(model, extensions, current_working_dir(), tls_config) + Self::from_env_with_working_dir(extensions, current_working_dir(), tls_config) } fn from_env_with_working_dir( - model: ModelConfig, extensions: Vec, working_dir: PathBuf, _tls_config: Option, @@ -103,12 +100,14 @@ impl ProviderDef for CodexAcpProvider { mcp_servers, // Disabled until https://github.com/zed-industries/codex-acp/issues/179 is fixed. session_mode_id: None, + session_config_options: vec![], + model_config_option_id: None, mode_mapping, notification_callback: None, }; let metadata = Self::metadata(); - AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await + AcpProvider::connect(metadata.name, goose_mode, provider_config).await }) } } diff --git a/crates/goose/src/providers/copilot_acp.rs b/crates/goose/src/providers/copilot_acp.rs index a2f732aa41d4..8685994ed086 100644 --- a/crates/goose/src/providers/copilot_acp.rs +++ b/crates/goose/src/providers/copilot_acp.rs @@ -11,7 +11,6 @@ use crate::config::{Config, GooseMode}; use crate::providers::base::{ current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, }; -use goose_providers::model::ModelConfig; pub(crate) const COPILOT_ACP_PROVIDER_NAME: &str = "copilot-acp"; const COPILOT_ACP_DOC_URL: &str = "https://github.com/github/copilot-cli"; @@ -46,15 +45,13 @@ impl ProviderDef for CopilotAcpProvider { type Provider = AcpProvider; fn from_env( - model: ModelConfig, extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Self::from_env_with_working_dir(model, extensions, current_working_dir(), tls_config) + Self::from_env_with_working_dir(extensions, current_working_dir(), tls_config) } fn from_env_with_working_dir( - model: ModelConfig, extensions: Vec, working_dir: PathBuf, _tls_config: Option, @@ -66,12 +63,16 @@ impl ProviderDef for CopilotAcpProvider { .with_npm() .resolve(COPILOT_ACP_BINARY)?; let goose_mode = config.get_goose_mode().unwrap_or(GooseMode::Auto); + let model = config + .get_goose_model() + .unwrap_or_else(|_| ACP_CURRENT_MODEL.to_string()); - let mut args = vec!["--acp".to_string()]; - if model.model_name != ACP_CURRENT_MODEL { - args.push("--model".to_string()); - args.push(model.model_name.clone()); - } + let args = vec!["--acp".to_string()]; + let session_config_options = if model == ACP_CURRENT_MODEL { + vec![] + } else { + vec![("model".to_string(), model)] + }; // Copilot modes are full protocol URIs. // No approve-specific mode; permissions are handled separately. @@ -90,12 +91,14 @@ impl ProviderDef for CopilotAcpProvider { work_dir: working_dir, mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), + session_config_options, + model_config_option_id: Some("model".to_string()), mode_mapping, notification_callback: None, }; let metadata = Self::metadata(); - AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await + AcpProvider::connect(metadata.name, goose_mode, provider_config).await }) } } diff --git a/crates/goose/src/providers/cursor_agent.rs b/crates/goose/src/providers/cursor_agent.rs index d275d44356d5..b36440fcff49 100644 --- a/crates/goose/src/providers/cursor_agent.rs +++ b/crates/goose/src/providers/cursor_agent.rs @@ -30,14 +30,12 @@ pub const CURSOR_AGENT_DOC_URL: &str = "https://docs.cursor.com/en/cli/overview" #[derive(Debug, serde::Serialize)] pub struct CursorAgentProvider { command: PathBuf, - model: ModelConfig, #[serde(skip)] name: String, } impl CursorAgentProvider { pub async fn from_env( - model: ModelConfig, _tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -46,7 +44,6 @@ impl CursorAgentProvider { Ok(Self { command: resolved_command, - model, name: CURSOR_AGENT_PROVIDER_NAME.to_string(), }) } @@ -182,6 +179,7 @@ impl CursorAgentProvider { async fn execute_command( &self, + model: &ModelConfig, system: &str, messages: &[Message], _tools: &[Tool], @@ -197,7 +195,7 @@ impl CursorAgentProvider { filter_extensions_from_system_prompt(system).len() ); println!("Full prompt: {}", prompt); - println!("Model: {}", self.model.model_name); + println!("Model: {}", model.model_name); println!("================================"); } @@ -208,7 +206,7 @@ impl CursorAgentProvider { cmd.env("PATH", path); } - cmd.arg("--model").arg(&self.model.model_name); + cmd.arg("--model").arg(&model.model_name); cmd.arg("-p") .arg(&prompt) @@ -303,11 +301,10 @@ impl ProviderDef for CursorAgentProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -317,11 +314,6 @@ impl Provider for CursorAgentProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - // Return the model config with appropriate context limit for Cursor models - self.model.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { Ok(CURSOR_AGENT_KNOWN_MODELS .iter() @@ -345,7 +337,9 @@ impl Provider for CursorAgentProvider { return Ok(stream_from_single_message(message, provider_usage)); } - let lines = self.execute_command(system, messages, tools).await?; + let lines = self + .execute_command(model_config, system, messages, tools) + .await?; let (message, usage) = self.parse_cursor_agent_response(&lines)?; @@ -362,7 +356,7 @@ impl Provider for CursorAgentProvider { "usage": usage }); - let mut log = start_log(&self.model, &payload)?; + let mut log = start_log(model_config, &payload)?; log.write(&response, Some(&usage))?; let provider_usage = ProviderUsage::new(model_config.model_name.clone(), usage); diff --git a/crates/goose/src/providers/databricks.rs b/crates/goose/src/providers/databricks.rs index 012a39d4b2a5..b9d1c0126185 100644 --- a/crates/goose/src/providers/databricks.rs +++ b/crates/goose/src/providers/databricks.rs @@ -63,7 +63,7 @@ static DATABRICKS_ENDPOINT_INFO_CACHE: LazyLock< Mutex>, > = LazyLock::new(|| Mutex::new(std::collections::HashMap::new())); pub const DATABRICKS_DEFAULT_MODEL: &str = "databricks-claude-sonnet-4"; -const DATABRICKS_DEFAULT_FAST_MODEL: &str = "databricks-claude-haiku-4-5"; +pub const DATABRICKS_DEFAULT_FAST_MODEL: &str = "databricks-claude-haiku-4-5"; pub const DATABRICKS_KNOWN_MODELS: &[&str] = &[ "databricks-claude-sonnet-4-5", "databricks-meta-llama-3-3-70b-instruct", @@ -80,7 +80,6 @@ pub struct DatabricksProvider { #[serde(skip)] host: String, auth: DatabricksAuth, - model: ModelConfig, image_format: ImageFormat, #[serde(skip)] retry_config: RetryConfig, @@ -98,7 +97,6 @@ impl DatabricksProvider { } pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -141,23 +139,16 @@ impl DatabricksProvider { tls_config.clone(), )?; - let mut provider = Self { + Ok(Self { api_client, host, auth, - model: model.clone(), image_format: ImageFormat::OpenAi, retry_config, name: DATABRICKS_PROVIDER_NAME.to_string(), token_cache, instance_id: Self::resolve_instance_id(), - }; - provider.model = crate::model_config::with_configured_fast_model( - model, - DATABRICKS_PROVIDER_NAME, - DATABRICKS_DEFAULT_FAST_MODEL, - )?; - Ok(provider) + }) } fn load_retry_config(config: &crate::config::Config) -> RetryConfig { @@ -487,12 +478,12 @@ impl DatabricksProvider { fn model_info_from_endpoint(info: DatabricksEndpointInfo) -> ModelInfo { let context_model = info.upstream_model_name.as_deref().unwrap_or(&info.name); - let context_limit = ModelConfig::new_or_fail(context_model) + let context_limit = ModelConfig::new(context_model) .with_canonical_limits(DATABRICKS_PROVIDER_NAME) .context_limit(); let reasoning = info .reasoning - .unwrap_or_else(|| ModelConfig::new_or_fail(context_model).is_reasoning_model()); + .unwrap_or_else(|| ModelConfig::new(context_model).is_reasoning_model()); ModelInfo { name: info.name, @@ -539,6 +530,7 @@ impl goose_providers::base::ProviderDescriptor for DatabricksProvider { ConfigKey::new("DATABRICKS_TOKEN", false, true, None, true), ], ) + .with_fast_model(DATABRICKS_DEFAULT_FAST_MODEL) } } @@ -546,11 +538,10 @@ impl ProviderDef for DatabricksProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -571,10 +562,6 @@ impl Provider for DatabricksProvider { Ok(()) } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn stream( &self, model_config: &ModelConfig, @@ -812,7 +799,10 @@ impl Provider for DatabricksProvider { Ok(Self::model_info_from_endpoint(endpoint_info)) } - async fn fetch_recommended_model_info(&self) -> Result, ProviderError> { + async fn fetch_recommended_model_info( + &self, + _toolshim: bool, + ) -> Result, ProviderError> { self.fetch_supported_model_info().await } } diff --git a/crates/goose/src/providers/databricks_v2.rs b/crates/goose/src/providers/databricks_v2.rs index 2c1d4b30ccb5..0d0ebd8aae47 100644 --- a/crates/goose/src/providers/databricks_v2.rs +++ b/crates/goose/src/providers/databricks_v2.rs @@ -54,7 +54,6 @@ enum DatabricksV2Route { pub struct DatabricksV2Provider { #[serde(skip)] api_client: ApiClient, - model: ModelConfig, #[serde(skip)] retry_config: RetryConfig, #[serde(skip)] @@ -69,7 +68,6 @@ impl DatabricksV2Provider { } pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -95,13 +93,12 @@ impl DatabricksV2Provider { DatabricksAuth::oauth(host.clone()) }; - Self::new(host, auth, model, retry_config, tls_config) + Self::new(host, auth, retry_config, tls_config) } fn new( host: String, auth: DatabricksAuth, - model: ModelConfig, retry_config: RetryConfig, tls_config: Option, ) -> Result { @@ -124,7 +121,6 @@ impl DatabricksV2Provider { Ok(Self { api_client, - model, retry_config, name: DATABRICKS_V2_PROVIDER_NAME.to_string(), token_cache, @@ -355,11 +351,10 @@ impl ProviderDef for DatabricksV2Provider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -379,10 +374,6 @@ impl Provider for DatabricksV2Provider { Ok(()) } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn stream( &self, model_config: &ModelConfig, diff --git a/crates/goose/src/providers/formats/anthropic.rs b/crates/goose/src/providers/formats/anthropic.rs index fa5315b36b43..c9018f8247ea 100644 --- a/crates/goose/src/providers/formats/anthropic.rs +++ b/crates/goose/src/providers/formats/anthropic.rs @@ -1595,20 +1595,13 @@ mod tests { } fn cfg(name: &str) -> ModelConfig { - ModelConfig { - model_name: name.to_string(), - ..Default::default() - } + ModelConfig::new(name) } fn cfg_with_effort(name: &str, effort: &str) -> ModelConfig { let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), json!(effort)); - ModelConfig { - model_name: name.to_string(), - request_params: Some(params), - ..Default::default() - } + ModelConfig::new(name).with_merged_request_params(params) } #[test] diff --git a/crates/goose/src/providers/formats/bedrock.rs b/crates/goose/src/providers/formats/bedrock.rs index 57f0eae340b1..08a225c9cbd7 100644 --- a/crates/goose/src/providers/formats/bedrock.rs +++ b/crates/goose/src/providers/formats/bedrock.rs @@ -535,7 +535,7 @@ mod tests { fn test_bedrock_anthropic_thinking_fields_enabled() { let mut params = HashMap::new(); params.insert("thinking_effort".to_string(), json!("low")); - let mut config = ModelConfig::new_or_fail("us.anthropic.claude-3-7-sonnet-20250219-v1:0"); + let mut config = ModelConfig::new("us.anthropic.claude-3-7-sonnet-20250219-v1:0"); config.request_params = Some(params); config.reasoning = Some(true); @@ -553,7 +553,7 @@ mod tests { #[test] fn test_bedrock_anthropic_thinking_fields_disabled() { - let mut config = ModelConfig::new_or_fail("us.anthropic.claude-3-7-sonnet-20250219-v1:0"); + let mut config = ModelConfig::new("us.anthropic.claude-3-7-sonnet-20250219-v1:0"); config.reasoning = Some(true); config.request_params = Some(HashMap::from([( "thinking_effort".to_string(), @@ -565,7 +565,7 @@ mod tests { #[test] fn test_bedrock_anthropic_thinking_fields_always_on_adaptive() { - let mut config = ModelConfig::new_or_fail("global.anthropic.claude-fable-5"); + let mut config = ModelConfig::new("global.anthropic.claude-fable-5"); config.reasoning = Some(true); config.request_params = Some(HashMap::from([( "thinking_effort".to_string(), @@ -584,7 +584,7 @@ mod tests { #[test] fn test_bedrock_anthropic_thinking_fields_adaptive_with_effort() { - let mut config = ModelConfig::new_or_fail("us.anthropic.claude-opus-4.7"); + let mut config = ModelConfig::new("us.anthropic.claude-opus-4.7"); config.reasoning = Some(true); config.request_params = Some(HashMap::from([( "thinking_effort".to_string(), @@ -603,7 +603,7 @@ mod tests { #[test] fn test_bedrock_anthropic_thinking_fields_adaptive_with_version_suffix() { - let mut config = ModelConfig::new_or_fail("us.anthropic.claude-opus-4-7-20251101-v1:0"); + let mut config = ModelConfig::new("us.anthropic.claude-opus-4-7-20251101-v1:0"); config.reasoning = Some(true); config.request_params = Some(HashMap::from([( "thinking_effort".to_string(), @@ -622,7 +622,7 @@ mod tests { #[test] fn test_bedrock_thinking_fields_skipped_for_non_anthropic() { - let mut config = ModelConfig::new_or_fail("us.deepseek.r1-v1:0"); + let mut config = ModelConfig::new("us.deepseek.r1-v1:0"); config.reasoning = Some(true); config.request_params = Some(HashMap::from([( "thinking_effort".to_string(), diff --git a/crates/goose/src/providers/formats/databricks.rs b/crates/goose/src/providers/formats/databricks.rs index b86542c87eab..e4b344a81dcb 100644 --- a/crates/goose/src/providers/formats/databricks.rs +++ b/crates/goose/src/providers/formats/databricks.rs @@ -1065,7 +1065,6 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1100,7 +1099,6 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: Some(params), reasoning: None, }; @@ -1120,7 +1118,6 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: Some(params), reasoning: None, }; @@ -1141,7 +1138,6 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: Some(params), reasoning: None, }; @@ -1160,7 +1156,6 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1179,7 +1174,6 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1198,7 +1192,6 @@ mod tests { max_tokens: Some(1024), toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1210,7 +1203,7 @@ mod tests { #[test] fn test_create_request_adaptive_thinking_for_46_models() -> anyhow::Result<()> { - let mut model_config = ModelConfig::new_or_fail("databricks-claude-opus-4-6"); + let mut model_config = ModelConfig::new("databricks-claude-opus-4-6"); model_config.max_tokens = Some(4096); let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), serde_json::json!("low")); @@ -1237,7 +1230,7 @@ mod tests { "databricks-claude-fable-5", "global.anthropic.claude-fable-5", ] { - let mut model_config = ModelConfig::new_or_fail(name); + let mut model_config = ModelConfig::new(name); model_config.max_tokens = Some(4096); let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), serde_json::json!("high")); @@ -1258,7 +1251,7 @@ mod tests { fn test_create_request_always_on_adaptive_off_effort_falls_back_to_high() -> anyhow::Result<()> { let _guard = env_lock::lock_env([("GOOSE_THINKING_EFFORT", None::<&str>)]); - let mut model_config = ModelConfig::new_or_fail("databricks-claude-fable-5"); + let mut model_config = ModelConfig::new("databricks-claude-fable-5"); model_config.max_tokens = Some(4096); let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), serde_json::json!("off")); @@ -1274,7 +1267,7 @@ mod tests { #[test] fn test_create_request_enabled_thinking_with_budget() -> anyhow::Result<()> { - let mut model_config = ModelConfig::new_or_fail("databricks-claude-3-7-sonnet"); + let mut model_config = ModelConfig::new("databricks-claude-3-7-sonnet"); model_config.max_tokens = Some(4096); let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), serde_json::json!("high")); @@ -1299,7 +1292,7 @@ mod tests { ("high", 16000), ("max", 32000), ] { - let mut model_config = ModelConfig::new_or_fail("databricks-claude-3-7-sonnet"); + let mut model_config = ModelConfig::new("databricks-claude-3-7-sonnet"); model_config.max_tokens = Some(4096); let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), serde_json::json!(effort)); @@ -1643,7 +1636,6 @@ mod tests { max_tokens: Some(8192), toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; @@ -1696,7 +1688,6 @@ mod tests { max_tokens: Some(4096), toolshim: false, toolshim_model: None, - fast_model_config: None, request_params: None, reasoning: None, }; diff --git a/crates/goose/src/providers/formats/google.rs b/crates/goose/src/providers/formats/google.rs index cf4d120cacd5..3c3387b970fa 100644 --- a/crates/goose/src/providers/formats/google.rs +++ b/crates/goose/src/providers/formats/google.rs @@ -1404,7 +1404,7 @@ data: [DONE]"#; // Test 1: Gemini 3 model with low thinking effort let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), serde_json::json!("low")); - let mut config = ModelConfig::new("gemini-3-pro").unwrap(); + let mut config = ModelConfig::new("gemini-3-pro"); config.request_params = Some(params); let result = get_thinking_config(&config); assert!(result.is_some()); @@ -1416,7 +1416,7 @@ data: [DONE]"#; // Test 2: Gemini 3 model with high thinking effort let mut params = std::collections::HashMap::new(); params.insert("thinking_effort".to_string(), serde_json::json!("high")); - let mut config = ModelConfig::new("Gemini-3-Flash").unwrap(); + let mut config = ModelConfig::new("Gemini-3-Flash"); config.request_params = Some(params); let result = get_thinking_config(&config); assert!(result.is_some()); @@ -1426,7 +1426,7 @@ data: [DONE]"#; Some(ThinkingLevel::High) )); - let config = ModelConfig::new("gemini-2.5-flash").unwrap(); + let config = ModelConfig::new("gemini-2.5-flash"); let result = get_thinking_config(&config); assert!(result.is_some()); let thinking_config = result.unwrap(); @@ -1439,9 +1439,7 @@ data: [DONE]"#; let mut params = HashMap::new(); params.insert("thinking_budget".to_string(), json!(4096)); - let config = ModelConfig::new("gemini-2.5-flash") - .unwrap() - .with_merged_request_params(params); + let config = ModelConfig::new("gemini-2.5-flash").with_merged_request_params(params); let result = get_thinking_config(&config); assert!(result.is_some()); let thinking_config = result.unwrap(); @@ -1449,9 +1447,7 @@ data: [DONE]"#; let mut params = HashMap::new(); params.insert("thinking_budget".to_string(), json!(-1)); - let config = ModelConfig::new("gemini-2.5-flash") - .unwrap() - .with_merged_request_params(params); + let config = ModelConfig::new("gemini-2.5-flash").with_merged_request_params(params); let result = get_thinking_config(&config); assert!(result.is_some()); let thinking_config = result.unwrap(); @@ -1460,11 +1456,11 @@ data: [DONE]"#; Some(GEMINI25_DEFAULT_THINKING_BUDGET) ); - let config = ModelConfig::new("gemini-2.0-flash").unwrap(); + let config = ModelConfig::new("gemini-2.0-flash"); let result = get_thinking_config(&config); assert!(result.is_none()); - let config = ModelConfig::new("gpt-4o").unwrap(); + let config = ModelConfig::new("gpt-4o"); let result = get_thinking_config(&config); assert!(result.is_none()); } diff --git a/crates/goose/src/providers/formats/openrouter.rs b/crates/goose/src/providers/formats/openrouter.rs index a68465fe9bc5..20cbda455751 100644 --- a/crates/goose/src/providers/formats/openrouter.rs +++ b/crates/goose/src/providers/formats/openrouter.rs @@ -190,7 +190,7 @@ mod tests { "messages": [], "reasoning_effort": "high" }); - let mut model_config = ModelConfig::new_or_fail("openai/gpt-5"); + let mut model_config = ModelConfig::new("openai/gpt-5"); let mut params = HashMap::new(); params.insert("thinking_effort".to_string(), json!("max")); model_config.request_params = Some(params); @@ -207,7 +207,7 @@ mod tests { "model": "x-ai/grok-4", "messages": [] }); - let mut model_config = ModelConfig::new_or_fail("x-ai/grok-4"); + let mut model_config = ModelConfig::new("x-ai/grok-4"); let mut params = HashMap::new(); params.insert("thinking_effort".to_string(), json!("high")); model_config.request_params = Some(params); @@ -224,7 +224,7 @@ mod tests { "model": "anthropic/claude-sonnet-4", "messages": [] }); - let mut model_config = ModelConfig::new_or_fail("anthropic/claude-sonnet-4"); + let mut model_config = ModelConfig::new("anthropic/claude-sonnet-4"); let mut params = HashMap::new(); params.insert("thinking_effort".to_string(), json!("high")); model_config.request_params = Some(params); @@ -240,7 +240,7 @@ mod tests { "model": "openai/gpt-4o", "messages": [] }); - let mut model_config = ModelConfig::new_or_fail("openai/gpt-4o"); + let mut model_config = ModelConfig::new("openai/gpt-4o"); let mut params = HashMap::new(); params.insert("thinking_effort".to_string(), json!("high")); model_config.request_params = Some(params); @@ -257,7 +257,7 @@ mod tests { "model": "x-ai/grok-4", "messages": [] }); - let mut model_config = ModelConfig::new_or_fail("x-ai/grok-4"); + let mut model_config = ModelConfig::new("x-ai/grok-4"); let mut params = HashMap::new(); params.insert("thinking_effort".to_string(), json!("off")); model_config.request_params = Some(params); diff --git a/crates/goose/src/providers/formats/snowflake.rs b/crates/goose/src/providers/formats/snowflake.rs index ba574f9ff8e7..bd72fa6e5845 100644 --- a/crates/goose/src/providers/formats/snowflake.rs +++ b/crates/goose/src/providers/formats/snowflake.rs @@ -562,8 +562,7 @@ data: {"id":"a9537c2c-2017-4906-9817-2456168d89fa","model":"claude-sonnet-4-2025 use crate::conversation::message::Message; use goose_providers::model::ModelConfig; - let model_config = - ModelConfig::new_or_fail("claude-4-sonnet").with_canonical_limits("snowflake"); + let model_config = ModelConfig::new("claude-4-sonnet").with_canonical_limits("snowflake"); let system = "You are a helpful assistant that can use tools to get information."; let messages = vec![Message::user().with_text("What is the stock price of Nvidia?")]; @@ -672,8 +671,7 @@ data: {"id":"a9537c2c-2017-4906-9817-2456168d89fa","model":"claude-sonnet-4-2025 use crate::conversation::message::Message; use goose_providers::model::ModelConfig; - let model_config = - ModelConfig::new_or_fail("claude-4-sonnet").with_canonical_limits("snowflake"); + let model_config = ModelConfig::new("claude-4-sonnet").with_canonical_limits("snowflake"); let system = "Reply with only a description in four words or less"; let messages = vec![Message::user().with_text("Test message")]; let tools = vec![Tool::new( diff --git a/crates/goose/src/providers/gcpvertexai.rs b/crates/goose/src/providers/gcpvertexai.rs index 88a3dae6389b..0b6240f0ff64 100644 --- a/crates/goose/src/providers/gcpvertexai.rs +++ b/crates/goose/src/providers/gcpvertexai.rs @@ -148,8 +148,6 @@ pub struct GcpVertexAIProvider { project_id: String, /// GCP region for model deployment location: String, - /// Configuration for the specific model being used - model: ModelConfig, /// Retry configuration for handling rate limit errors #[serde(skip)] retry_config: RetryConfig, @@ -166,7 +164,6 @@ impl GcpVertexAIProvider { /// # Arguments /// * `model` - Configuration for the model to be used pub async fn from_env( - model: ModelConfig, _tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -189,7 +186,6 @@ impl GcpVertexAIProvider { host, project_id, location, - model, retry_config, name: GCP_VERTEX_AI_PROVIDER_NAME.to_string(), }) @@ -262,6 +258,7 @@ impl GcpVertexAIProvider { fn build_request_url( &self, + model: &ModelConfig, provider: ModelProvider, location: &str, streaming: bool, @@ -270,7 +267,7 @@ impl GcpVertexAIProvider { &self.host, &self.location, &self.project_id, - &self.model.model_name, + &model.model_name, provider, location, streaming, @@ -395,13 +392,14 @@ impl GcpVertexAIProvider { async fn post_stream_with_location( &self, + model: &ModelConfig, session_id: Option<&str>, payload: &Value, context: &RequestContext, location: &str, ) -> Result { let url = self - .build_request_url(context.provider(), location, true) + .build_request_url(model, context.provider(), location, true) .map_err(|e| ProviderError::RequestFailed(e.to_string()))?; self.send_request_with_retry(session_id, url, payload).await @@ -409,12 +407,13 @@ impl GcpVertexAIProvider { async fn post_stream( &self, + model: &ModelConfig, session_id: Option<&str>, payload: &Value, context: &RequestContext, ) -> Result { let result = self - .post_stream_with_location(session_id, payload, context, &self.location) + .post_stream_with_location(model, session_id, payload, context, &self.location) .await; if self.location == context.model.known_location().to_string() || result.is_ok() { @@ -431,7 +430,7 @@ impl GcpVertexAIProvider { "Trying known location {known_location} for {model_name} instead of {configured_location}: {msg}" ); - self.post_stream_with_location(session_id, payload, context, &known_location) + self.post_stream_with_location(model, session_id, payload, context, &known_location) .await } _ => result, @@ -583,11 +582,10 @@ impl ProviderDef for GcpVertexAIProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -603,11 +601,6 @@ impl Provider for GcpVertexAIProvider { /// * `system` - System prompt or context /// * `messages` - Array of previous messages in the conversation /// * `tools` - Array of available tools for the model - /// Returns the current model configuration. - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn stream( &self, model_config: &ModelConfig, @@ -627,7 +620,7 @@ impl Provider for GcpVertexAIProvider { let mut log = start_log(model_config, &request)?; let response = self - .post_stream(Some(session_id), &request, &context) + .post_stream(model_config, Some(session_id), &request, &context) .await .inspect_err(|e| { let _ = log.error(e); diff --git a/crates/goose/src/providers/gemini_cli.rs b/crates/goose/src/providers/gemini_cli.rs index d834b80a4c37..665d91081df5 100644 --- a/crates/goose/src/providers/gemini_cli.rs +++ b/crates/goose/src/providers/gemini_cli.rs @@ -38,7 +38,6 @@ pub const GEMINI_CLI_DOC_URL: &str = "https://ai.google.dev/gemini-api/docs"; #[derive(Debug, serde::Serialize)] pub struct GeminiCliProvider { command: PathBuf, - model: ModelConfig, #[serde(skip)] name: String, #[serde(skip)] @@ -47,7 +46,6 @@ pub struct GeminiCliProvider { impl GeminiCliProvider { pub async fn from_env( - model: ModelConfig, _tls_config: Option, ) -> Result { let config = Config::global(); @@ -56,7 +54,6 @@ impl GeminiCliProvider { Ok(Self { command: resolved_command, - model, name: GEMINI_CLI_PROVIDER_NAME.to_string(), cli_session_id: Arc::new(OnceLock::new()), }) @@ -180,11 +177,10 @@ impl ProviderDef for GeminiCliProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -198,10 +194,6 @@ impl Provider for GeminiCliProvider { true } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { Ok(GEMINI_CLI_KNOWN_MODELS .iter() @@ -336,7 +328,6 @@ mod tests { fn make_provider() -> GeminiCliProvider { GeminiCliProvider { command: PathBuf::from("gemini"), - model: ModelConfig::new("gemini-2.5-pro").unwrap(), name: "gemini-cli".to_string(), cli_session_id: Arc::new(OnceLock::new()), } diff --git a/crates/goose/src/providers/gemini_oauth.rs b/crates/goose/src/providers/gemini_oauth.rs index 92b24cc2452e..dfe5ce720de0 100644 --- a/crates/goose/src/providers/gemini_oauth.rs +++ b/crates/goose/src/providers/gemini_oauth.rs @@ -831,29 +831,20 @@ fn parse_retry_delay(body: &str) -> Option { pub struct GeminiOAuthProvider { #[serde(skip)] token_provider: Arc, - model: ModelConfig, #[serde(skip)] name: String, } impl GeminiOAuthProvider { pub async fn from_env( - model: ModelConfig, _tls_config: Option, ) -> Result { - let model = crate::model_config::with_configured_fast_model( - model, - GEMINI_OAUTH_PROVIDER_NAME, - GEMINI_OAUTH_DEFAULT_FAST_MODEL, - )?; - let token_provider = Arc::new(GeminiOAuthTokenProvider::new( GeminiOAuthAuthState::instance(), )); Ok(Self { token_provider, - model, name: GEMINI_OAUTH_PROVIDER_NAME.to_string(), }) } @@ -952,6 +943,7 @@ impl goose_providers::base::ProviderDescriptor for GeminiOAuthProvider { false, )], ) + .with_fast_model(GEMINI_OAUTH_DEFAULT_FAST_MODEL) } } @@ -959,11 +951,10 @@ impl ProviderDef for GeminiOAuthProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -973,10 +964,6 @@ impl Provider for GeminiOAuthProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn configure_oauth(&self) -> Result<(), ProviderError> { self.token_provider .get_valid_setup() diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index a9c2024a1e2b..8c3e9fce60f9 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -195,7 +195,6 @@ pub struct GithubCopilotProvider { cache: DiskCache, #[serde(skip)] mu: tokio::sync::Mutex>>, - model: ModelConfig, #[serde(skip)] urls: GithubCopilotUrls, #[serde(skip)] @@ -232,7 +231,6 @@ impl GithubCopilotProvider { } pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { let config = Config::global(); @@ -255,7 +253,6 @@ impl GithubCopilotProvider { client, cache, mu, - model, urls, client_id, name: GITHUB_COPILOT_PROVIDER_NAME.to_string(), @@ -553,11 +550,10 @@ impl ProviderDef for GithubCopilotProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -567,10 +563,6 @@ impl Provider for GithubCopilotProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - #[tracing::instrument( skip(self, model_config, session_id, system, messages, tools), fields(session.id = %session_id, gen_ai.request.model = %model_config.model_name) diff --git a/crates/goose/src/providers/google.rs b/crates/goose/src/providers/google.rs index 6d2be4b78ffb..3ab13a819fc2 100644 --- a/crates/goose/src/providers/google.rs +++ b/crates/goose/src/providers/google.rs @@ -60,22 +60,14 @@ pub const GOOGLE_DOC_URL: &str = "https://ai.google.dev/gemini-api/docs/models"; pub struct GoogleProvider { #[serde(skip)] api_client: ApiClient, - model: ModelConfig, #[serde(skip)] name: String, } impl GoogleProvider { pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { - let model = crate::model_config::with_configured_fast_model( - model, - GOOGLE_PROVIDER_NAME, - GOOGLE_DEFAULT_FAST_MODEL, - )?; - let config = crate::config::Config::global(); let api_key: String = config.get_secret("GOOGLE_API_KEY")?; let host: String = config @@ -92,7 +84,6 @@ impl GoogleProvider { Ok(Self { api_client, - model, name: GOOGLE_PROVIDER_NAME.to_string(), }) } @@ -126,6 +117,7 @@ impl goose_providers::base::ProviderDescriptor for GoogleProvider { ConfigKey::new("GOOGLE_HOST", false, false, Some(GOOGLE_API_HOST), false), ], ) + .with_fast_model(GOOGLE_DEFAULT_FAST_MODEL) .with_setup_steps(vec![ "Go to https://aistudio.google.com and sign in with your Google account", "Click 'Get API key' on the left sidebar", @@ -139,11 +131,10 @@ impl ProviderDef for GoogleProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -153,10 +144,6 @@ impl Provider for GoogleProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client diff --git a/crates/goose/src/providers/huggingface.rs b/crates/goose/src/providers/huggingface.rs index 2d071ede9a22..b24989fdedcc 100644 --- a/crates/goose/src/providers/huggingface.rs +++ b/crates/goose/src/providers/huggingface.rs @@ -72,7 +72,6 @@ impl HuggingFaceProvider { } pub fn from_custom_config( - model: ModelConfig, config: DeclarativeProviderConfig, tls_config: Option, ) -> Result { @@ -110,17 +109,10 @@ impl HuggingFaceProvider { api_client = api_client.with_headers(header_map)?; } - let model = if let Some(ref fast_model_name) = config.fast_model { - crate::model_config::with_configured_fast_model(model, &config.name, fast_model_name)? - } else { - model - }; - Ok(Self { inner: OpenAiCompatibleProvider::new( config.name.clone(), api_client, - model, completions_prefix, ) .with_supports_streaming(config.supports_streaming.unwrap_or(true)), @@ -140,10 +132,6 @@ impl Provider for HuggingFaceProvider { self.inner.get_name() } - fn get_model_config(&self) -> ModelConfig { - self.inner.get_model_config() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { if let Some(custom_models) = &self.custom_models { if self.dynamic_models == Some(false) { @@ -208,7 +196,6 @@ impl ProviderDef for HuggingFaceProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { @@ -225,7 +212,6 @@ impl ProviderDef for HuggingFaceProvider { inner: OpenAiCompatibleProvider::new( huggingface_auth::HUGGINGFACE_PROVIDER_NAME.to_string(), api_client, - model, String::new(), ), custom_models: None, @@ -449,12 +435,7 @@ mod tests { ModelInfo::new("static-b".to_string(), 128000), ]; - let provider = HuggingFaceProvider::from_custom_config( - ModelConfig::new("static-a").unwrap(), - config, - None, - ) - .unwrap(); + let provider = HuggingFaceProvider::from_custom_config(config, None).unwrap(); assert_eq!( provider.fetch_supported_models().await.unwrap(), @@ -468,11 +449,7 @@ mod tests { config.requires_auth = false; config.dynamic_models = Some(false); - let error = match HuggingFaceProvider::from_custom_config( - ModelConfig::new("model").unwrap(), - config, - None, - ) { + let error = match HuggingFaceProvider::from_custom_config(config, None) { Ok(_) => panic!("expected dynamic_models: false without static models to fail"), Err(error) => error, }; diff --git a/crates/goose/src/providers/init.rs b/crates/goose/src/providers/init.rs index e42e05822924..4861fe7a932b 100644 --- a/crates/goose/src/providers/init.rs +++ b/crates/goose/src/providers/init.rs @@ -48,7 +48,6 @@ use crate::{ providers::provider_registry::ProviderEntry, }; use anyhow::Result; -use goose_providers::model::ModelConfig; use tokio::sync::OnceCell; static REGISTRY: OnceCell> = OnceCell::const_new(); @@ -221,25 +220,18 @@ pub async fn inventory_identity(name: &str) -> Result, -) -> Result> { +pub async fn create(name: &str, extensions: Vec) -> Result> { let entry = get_from_registry(name).await?; - entry.create(model, extensions).await + entry.create(extensions).await } pub async fn create_with_working_dir( name: &str, - model: ModelConfig, extensions: Vec, working_dir: PathBuf, ) -> Result> { let entry = get_from_registry(name).await?; - entry - .create_with_working_dir(model, extensions, working_dir) - .await + entry.create_with_working_dir(extensions, working_dir).await } pub async fn create_with_default_model( @@ -268,11 +260,9 @@ pub async fn cleanup_provider(name: &str) -> Result<()> { pub async fn create_with_named_model( provider_name: &str, - model_name: &str, extensions: Vec, ) -> Result> { - let config = crate::model_config::model_config_from_user_config(provider_name, model_name)?; - create(provider_name, config, extensions).await + create(provider_name, extensions).await } #[cfg(test)] @@ -516,15 +506,27 @@ mod tests { .await .expect("custom providers should refresh"); - let provider = create_with_named_model("custom_inf", "kimi-k2.5", Vec::new()) + let inf_entry = get_from_registry("custom_inf") .await - .expect("custom_inf provider should be creatable"); - assert_eq!(provider.get_model_config().context_limit, Some(256_000)); - - let zero_provider = create_with_named_model("custom_zero", "zero-model", Vec::new()) + .expect("custom_inf entry should exist"); + let inf_config = inf_entry + .normalize_model_config( + crate::model_config::model_config_from_user_config("custom_inf", "kimi-k2.5") + .expect("custom_inf model config should resolve"), + ) + .expect("custom_inf model config should normalize"); + assert_eq!(inf_config.context_limit, Some(256_000)); + + let zero_entry = get_from_registry("custom_zero") .await - .expect("custom_zero provider should be creatable"); - assert_eq!(zero_provider.get_model_config().context_limit, None); + .expect("custom_zero entry should exist"); + let zero_config = zero_entry + .normalize_model_config( + crate::model_config::model_config_from_user_config("custom_zero", "zero-model") + .expect("custom_zero model config should resolve"), + ) + .expect("custom_zero model config should normalize"); + assert_eq!(zero_config.context_limit, None); std::env::remove_var("GOOSE_PATH_ROOT"); } diff --git a/crates/goose/src/providers/inventory/mod.rs b/crates/goose/src/providers/inventory/mod.rs index 28d9ca73fe5e..5177799394a4 100644 --- a/crates/goose/src/providers/inventory/mod.rs +++ b/crates/goose/src/providers/inventory/mod.rs @@ -617,9 +617,11 @@ impl ProviderInventoryService { .await { Ok(()) => { - match AssertUnwindSafe(provider.fetch_recommended_models()) - .catch_unwind() - .await + match AssertUnwindSafe(provider.fetch_recommended_models( + crate::model_config::global_toolshim(), + )) + .catch_unwind() + .await { Ok(Ok(models)) => Ok(models), Ok(Err(error)) => Err(anyhow::anyhow!(error.to_string())), diff --git a/crates/goose/src/providers/kimicode.rs b/crates/goose/src/providers/kimicode.rs index 6781ddfd9b5c..fba0bd116639 100644 --- a/crates/goose/src/providers/kimicode.rs +++ b/crates/goose/src/providers/kimicode.rs @@ -35,7 +35,6 @@ use rmcp::model::Tool; const KIMI_CODE_PROVIDER_NAME: &str = "kimi_code"; pub const KIMI_CODE_DEFAULT_MODEL: &str = "kimi-for-coding"; -pub const KIMI_CODE_DEFAULT_FAST_MODEL: &str = "kimi-for-coding"; /// Known models for the provider metadata registration. The live catalogue is /// fetched from `/v1/models` at request time; this constant is only used for /// `ProviderMetadata`. As of 2025-10 Kimi Code exposes a single model, @@ -152,7 +151,6 @@ pub struct KimiCodeProvider { auth_host: String, #[serde(skip)] api_base: String, - model: ModelConfig, #[serde(skip)] name: String, } @@ -163,14 +161,8 @@ impl KimiCodeProvider { } pub async fn from_env( - model: ModelConfig, _tls_config: Option, ) -> Result { - let model = crate::model_config::with_configured_fast_model( - model, - KIMI_CODE_PROVIDER_NAME, - KIMI_CODE_DEFAULT_FAST_MODEL, - )?; let client = Client::builder() .timeout(StdDuration::from_secs(DEFAULT_PROVIDER_TIMEOUT_SECS)) .build()?; @@ -182,7 +174,6 @@ impl KimiCodeProvider { device_id, auth_host: KIMI_AUTH_HOST.to_string(), api_base: KIMI_API_BASE.to_string(), - model, name: KIMI_CODE_PROVIDER_NAME.to_string(), }) } @@ -375,11 +366,10 @@ impl ProviderDef for KimiCodeProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -391,10 +381,6 @@ impl Provider for KimiCodeProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn stream( &self, model_config: &ModelConfig, @@ -513,7 +499,6 @@ mod tests { device_id: device_id.to_string(), auth_host: server_uri.to_string(), api_base: server_uri.to_string(), - model: ModelConfig::new(KIMI_CODE_DEFAULT_MODEL).unwrap(), name: KIMI_CODE_PROVIDER_NAME.to_string(), } } diff --git a/crates/goose/src/providers/litellm.rs b/crates/goose/src/providers/litellm.rs index ed446ee29cdc..9fed2f93e2ac 100644 --- a/crates/goose/src/providers/litellm.rs +++ b/crates/goose/src/providers/litellm.rs @@ -29,7 +29,6 @@ pub struct LiteLLMProvider { #[serde(skip)] api_client: ApiClient, base_path: String, - model: ModelConfig, #[serde(skip)] name: String, #[serde(skip)] @@ -38,7 +37,6 @@ pub struct LiteLLMProvider { impl LiteLLMProvider { pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -86,7 +84,6 @@ impl LiteLLMProvider { Ok(Self { api_client, base_path, - model, name: LITELLM_PROVIDER_NAME.to_string(), cached_model_info: tokio::sync::OnceCell::new(), }) @@ -153,6 +150,16 @@ impl LiteLLMProvider { .await?; handle_response_openai_compat(response).await } + + async fn supports_cache_control(&self, model: &ModelConfig) -> bool { + if let Ok(models) = self.get_or_fetch_models().await { + if let Some(model_info) = models.iter().find(|m| m.name == model.model_name) { + return model_info.supports_cache_control.unwrap_or(false); + } + } + + model.model_name.to_lowercase().contains("claude") + } } impl goose_providers::base::ProviderDescriptor for LiteLLMProvider { @@ -191,11 +198,10 @@ impl ProviderDef for LiteLLMProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -205,23 +211,25 @@ impl Provider for LiteLLMProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - let mut config = self.model.clone(); + async fn get_context_limit(&self, model_config: &ModelConfig) -> Result { + if let Some(limit) = model_config.context_limit { + return Ok(limit); + } + // The cache is populated lazily by the first stream() call (via // supports_cache_control). On turn 1 this will be None and we fall // back to DEFAULT_CONTEXT_LIMIT, which is fine — the conversation is // too small to trigger compaction. From turn 2 onward the real limit // from /model/info is used. - if config.context_limit.is_none() { - if let Some(models) = self.cached_model_info.get() { - if let Some(info) = models.iter().find(|m| m.name == config.model_name) { - if info.context_limit > 0 { - config.context_limit = Some(info.context_limit); - } + if let Some(models) = self.cached_model_info.get() { + if let Some(info) = models.iter().find(|m| m.name == model_config.model_name) { + if info.context_limit > 0 { + return Ok(info.context_limit); } } } - config + + Ok(model_config.context_limit()) } async fn stream( @@ -246,7 +254,7 @@ impl Provider for LiteLLMProvider { false, )?; - if self.supports_cache_control().await { + if self.supports_cache_control(model_config).await { payload = update_request_for_cache_control(&payload); } @@ -269,16 +277,6 @@ impl Provider for LiteLLMProvider { )) } - async fn supports_cache_control(&self) -> bool { - if let Ok(models) = self.get_or_fetch_models().await { - if let Some(model_info) = models.iter().find(|m| m.name == self.model.model_name) { - return model_info.supports_cache_control.unwrap_or(false); - } - } - - self.model.model_name.to_lowercase().contains("claude") - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { let models = self.get_or_fetch_models().await?; Ok(models.iter().map(|m| m.name.clone()).collect()) diff --git a/crates/goose/src/providers/local_inference.rs b/crates/goose/src/providers/local_inference.rs index e0d131c505db..fc33f06d0b9f 100644 --- a/crates/goose/src/providers/local_inference.rs +++ b/crates/goose/src/providers/local_inference.rs @@ -478,16 +478,14 @@ type StreamSender = pub struct LocalInferenceProvider { runtime: Arc, - model_config: ModelConfig, name: String, } impl LocalInferenceProvider { - pub async fn from_env(model: ModelConfig, _extensions: Vec) -> Result { + pub async fn from_env(_extensions: Vec) -> Result { let runtime = InferenceRuntime::get_or_init()?; Ok(Self { runtime, - model_config: model, name: PROVIDER_NAME.to_string(), }) } @@ -532,14 +530,13 @@ impl ProviderDef for LocalInferenceProvider { type Provider = Self; fn from_env( - model: ModelConfig, extensions: Vec, _tls_config: Option, ) -> BoxFuture<'static, Result> where Self: Sized, { - Box::pin(Self::from_env(model, extensions)) + Box::pin(Self::from_env(extensions)) } } @@ -549,10 +546,6 @@ impl Provider for LocalInferenceProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model_config.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { use crate::providers::local_inference::local_model_registry::get_registry; @@ -654,7 +647,7 @@ impl Provider for LocalInferenceProvider { }, }); - let mut log = start_log(&self.model_config, &log_payload)?; + let mut log = start_log(model_config, &log_payload)?; let (tx, mut rx) = tokio::sync::mpsc::channel::< Result<(Option, Option), ProviderError>, diff --git a/crates/goose/src/providers/nanogpt.rs b/crates/goose/src/providers/nanogpt.rs index 3b9371667c71..03cee2b308b9 100644 --- a/crates/goose/src/providers/nanogpt.rs +++ b/crates/goose/src/providers/nanogpt.rs @@ -24,7 +24,6 @@ const NANOGPT_API_KEY: &str = "NANOGPT_API_KEY"; pub struct NanoGptProvider { #[serde(skip)] api_client: ApiClient, - model: ModelConfig, #[serde(skip)] name: String, } @@ -64,7 +63,6 @@ impl NanoGptProvider { } pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -83,7 +81,6 @@ impl NanoGptProvider { Ok(Self { api_client, - model, name: NANOGPT_PROVIDER_NAME.to_string(), }) } @@ -107,11 +104,10 @@ impl ProviderDef for NanoGptProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -121,10 +117,6 @@ impl Provider for NanoGptProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self .api_client diff --git a/crates/goose/src/providers/ollama.rs b/crates/goose/src/providers/ollama.rs index ca94704f5a7d..834e5ce72f30 100644 --- a/crates/goose/src/providers/ollama.rs +++ b/crates/goose/src/providers/ollama.rs @@ -52,7 +52,6 @@ const OLLAMA_MAX_RETRY_INTERVAL_MS: u64 = 15_000; pub struct OllamaProvider { #[serde(skip)] api_client: ApiClient, - model: ModelConfig, supports_streaming: bool, name: String, skip_canonical_filtering: bool, @@ -131,7 +130,6 @@ pub(crate) fn ollama_host_configured(config: &crate::config::Config) -> bool { impl OllamaProvider { pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -170,7 +168,6 @@ impl OllamaProvider { Ok(Self { api_client, - model, supports_streaming: true, name: OLLAMA_PROVIDER_NAME.to_string(), skip_canonical_filtering: false, @@ -178,7 +175,6 @@ impl OllamaProvider { } pub fn from_custom_config( - model: ModelConfig, config: DeclarativeProviderConfig, tls_config: Option, ) -> Result { @@ -230,15 +226,8 @@ impl OllamaProvider { )); } - let model = if let Some(ref fast_model_name) = config.fast_model { - crate::model_config::with_configured_fast_model(model, &config.name, fast_model_name)? - } else { - model - }; - Ok(Self { api_client, - model, supports_streaming, name: config.name.clone(), skip_canonical_filtering: config.skip_canonical_filtering, @@ -273,11 +262,10 @@ impl ProviderDef for OllamaProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -291,10 +279,6 @@ impl Provider for OllamaProvider { self.skip_canonical_filtering } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - fn retry_config(&self) -> RetryConfig { RetryConfig::new( OLLAMA_MAX_RETRIES, @@ -522,9 +506,7 @@ mod tests { #[test] fn test_apply_ollama_options_uses_input_limit() { let _guard = env_lock::lock_env([("GOOSE_INPUT_LIMIT", Some("8192"))]); - let model_config = ModelConfig::new("qwen3") - .unwrap() - .with_context_limit(Some(16_000)); + let model_config = ModelConfig::new("qwen3").with_context_limit(Some(16_000)); let mut payload = json!({}); apply_ollama_options(&mut payload, &model_config); assert_eq!(payload["options"]["num_ctx"], 8192); @@ -533,9 +515,7 @@ mod tests { #[test] fn test_apply_ollama_options_falls_back_to_context_limit() { let _guard = env_lock::lock_env([("GOOSE_INPUT_LIMIT", None::<&str>)]); - let model_config = ModelConfig::new("qwen3") - .unwrap() - .with_context_limit(Some(12_000)); + let model_config = ModelConfig::new("qwen3").with_context_limit(Some(12_000)); let mut payload = json!({}); apply_ollama_options(&mut payload, &model_config); assert_eq!(payload["options"]["num_ctx"], 12_000); @@ -544,7 +524,7 @@ mod tests { #[test] fn test_apply_ollama_options_skips_when_no_limit() { let _guard = env_lock::lock_env([("GOOSE_INPUT_LIMIT", None::<&str>)]); - let mut model_config = ModelConfig::new("qwen3").unwrap(); + let mut model_config = ModelConfig::new("qwen3"); model_config.context_limit = None; let mut payload = json!({}); apply_ollama_options(&mut payload, &model_config); @@ -555,9 +535,7 @@ mod tests { fn test_raw_create_request_contains_unsupported_ollama_fields() { use crate::providers::formats::ollama::create_request; - let model_config = ModelConfig::new("llama3.1") - .unwrap() - .with_max_tokens(Some(4096)); + let model_config = ModelConfig::new("llama3.1").with_max_tokens(Some(4096)); let messages = vec![crate::conversation::message::Message::user().with_text("hi")]; let payload = create_request( @@ -588,9 +566,7 @@ mod tests { ("GOOSE_INPUT_LIMIT", None::<&str>), ("OLLAMA_STREAM_USAGE", None::<&str>), ]); - let model_config = ModelConfig::new("llama3.1") - .unwrap() - .with_max_tokens(Some(4096)); + let model_config = ModelConfig::new("llama3.1").with_max_tokens(Some(4096)); let messages = vec![crate::conversation::message::Message::user().with_text("hi")]; let mut payload = create_request( @@ -632,9 +608,7 @@ mod tests { ("GOOSE_INPUT_LIMIT", None::<&str>), ("OLLAMA_STREAM_USAGE", Some("false")), ]); - let model_config = ModelConfig::new("llama3.1") - .unwrap() - .with_max_tokens(Some(4096)); + let model_config = ModelConfig::new("llama3.1").with_max_tokens(Some(4096)); let messages = vec![crate::conversation::message::Message::user().with_text("hi")]; let mut payload = create_request( diff --git a/crates/goose/src/providers/openai_def.rs b/crates/goose/src/providers/openai_def.rs index d65664103044..bf1615fdc21b 100644 --- a/crates/goose/src/providers/openai_def.rs +++ b/crates/goose/src/providers/openai_def.rs @@ -6,35 +6,46 @@ use std::collections::HashMap; use crate::config::declarative_providers::DeclarativeProviderConfig; use crate::providers::base::{ProviderDef, DEFAULT_PROVIDER_TIMEOUT_SECS}; use goose_providers::api_client::{ApiClient, AuthMethod}; -use goose_providers::model::ModelConfig; use goose_providers::openai::{ ensure_url_scheme, parse_custom_headers, parse_openai_base_url, OpenAiProvider, OpenAiProviderBuilder, OPEN_AI_DEFAULT_BASE_PATH, OPEN_AI_DEFAULT_FAST_MODEL, - OPEN_AI_PROVIDER_NAME, OPEN_AI_VERSIONLESS_BASE_PATH, + OPEN_AI_VERSIONLESS_BASE_PATH, }; pub struct OpenAiProviderDef; impl ProviderDescriptor for OpenAiProviderDef { fn metadata() -> goose_providers::base::ProviderMetadata { + // The default fast model is resolved live in `live_fast_model` rather + // than baked into metadata here: registry metadata is snapshotted at + // init time, but the OpenAI base URL can change at runtime (e.g. + // switching to an OpenAI-compatible endpoint), which would otherwise + // leave the cached fast model stale. OpenAiProvider::metadata() } } +pub fn live_fast_model() -> Option { + match resolve_base_url(crate::config::Config::global()) { + Ok(parsed) if is_direct_openai_host(&parsed.host) => { + Some(OPEN_AI_DEFAULT_FAST_MODEL.to_string()) + } + _ => None, + } +} + impl ProviderDef for OpenAiProviderDef { type Provider = OpenAiProvider; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(from_env(model, tls_config)) + Box::pin(from_env(tls_config)) } } pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -54,34 +65,7 @@ pub async fn from_env( // otherwise "chat/completions" to match the OpenAI SDK convention. // // OPENAI_BASE_PATH always wins when set explicitly. - let parsed = if let Ok(h) = std::env::var("OPENAI_HOST") { - // OPENAI_HOST env var takes priority as a session override so - // that existing scripts like `OPENAI_HOST=… goose` still work - // even after OPENAI_BASE_URL is persisted in config. - ParsedBaseUrl { - host: h, - query_params: vec![], - has_v1: true, - from_base_url: false, - } - } else if let Some(raw_url) = config - .get_param::("OPENAI_BASE_URL") - .ok() - .map(|s| s.trim().to_string()) - .filter(|s| !s.is_empty()) - { - parse_base_url(&raw_url)? - } else { - let h: String = config - .get_param("OPENAI_HOST") - .unwrap_or_else(|_| "https://api.openai.com".to_string()); - ParsedBaseUrl { - host: h, - query_params: vec![], - has_v1: true, - from_base_url: false, - } - }; + let parsed = resolve_base_url(config)?; // When the host was derived from OPENAI_BASE_URL, read // OPENAI_BASE_PATH from env only so that the desktop UI's persisted @@ -104,28 +88,7 @@ pub async fn from_env( .unwrap_or_else(|_| default_bp()) }; - // Only apply the default fast model when talking to OpenAI directly. - // Custom/compatible endpoints likely don't serve gpt-4o-mini, so - // leave fast_model unset (complete_fast will fall back to the main model). - // Parse the URL and compare the hostname exactly to avoid false positives - // (e.g. https://api.openai.com.local:8000 or proxy paths containing api.openai.com). - let host = parsed.host.clone(); - - let is_openai = url::Url::parse(&host) - .ok() - .and_then(|u| u.host_str().map(|h| h.to_ascii_lowercase())) - .map(|h| h == "api.openai.com" || h.ends_with(".api.openai.com")) - .unwrap_or(false); - let model = if is_openai { - crate::model_config::with_configured_fast_model( - model, - OPEN_AI_PROVIDER_NAME, - OPEN_AI_DEFAULT_FAST_MODEL, - )? - } else { - model - }; - + let is_openai = is_direct_openai_host(&parsed.host); let secrets = config .get_secrets("OPENAI_API_KEY", &["OPENAI_CUSTOM_HEADERS"]) .unwrap_or_default(); @@ -174,7 +137,7 @@ pub async fn from_env( api_client = api_client.with_headers(header_map)?; } - let mut provider = OpenAiProviderBuilder::new(api_client, model) + let provider = OpenAiProviderBuilder::new(api_client) .base_path(base_path) .organization(organization) .project(project) @@ -182,7 +145,8 @@ pub async fn from_env( .preserve_thinking_context(!is_openai) .build(); - provider.probe_context_limit_if_unset().await; + // TODO(jack): replace this + // provider.probe_context_limit_if_unset(&mut model).await; Ok(provider) } @@ -235,7 +199,6 @@ pub fn resolve_api_key( } pub fn from_custom_config( - model: ModelConfig, config: DeclarativeProviderConfig, tls_config: Option, ) -> Result { @@ -307,13 +270,7 @@ pub fn from_custom_config( api_client = api_client.with_headers(header_map)?; } - let model = if let Some(ref fast_model_name) = config.fast_model { - crate::model_config::with_configured_fast_model(model, &config.name, fast_model_name)? - } else { - model - }; - - Ok(OpenAiProviderBuilder::new(api_client, model) + Ok(OpenAiProviderBuilder::new(api_client) .base_path(base_path) .custom_headers(config.headers) .supports_streaming(config.supports_streaming.unwrap_or(true)) @@ -350,6 +307,56 @@ fn parse_base_url(raw_url: &str) -> Result { }) } +/// Resolve the effective host from environment and config. +/// +/// Priority (highest first): +/// 1. OPENAI_HOST env var — session override (deprecated but still honoured) +/// 2. OPENAI_BASE_URL (env or config) — ecosystem-standard +/// 3. OPENAI_HOST from config file — persisted by `goose configure` +/// 4. Default "https://api.openai.com" +fn resolve_base_url(config: &crate::config::Config) -> Result { + if let Ok(h) = std::env::var("OPENAI_HOST") { + return Ok(ParsedBaseUrl { + host: h, + query_params: vec![], + has_v1: true, + from_base_url: false, + }); + } + + if let Some(raw_url) = config + .get_param::("OPENAI_BASE_URL") + .ok() + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) + { + return parse_base_url(&raw_url); + } + + let h: String = config + .get_param("OPENAI_HOST") + .unwrap_or_else(|_| "https://api.openai.com".to_string()); + Ok(ParsedBaseUrl { + host: h, + query_params: vec![], + has_v1: true, + from_base_url: false, + }) +} + +/// Whether `host` points at OpenAI directly. +/// +/// Compares the hostname exactly to avoid false positives (e.g. +/// `https://api.openai.com.local:8000` or proxy paths containing +/// `api.openai.com`). +fn is_direct_openai_host(host: &str) -> bool { + url::Url::parse(host) + .ok() + .and_then(|u| u.host_str().map(|h| h.to_ascii_lowercase())) + .map(|h| h == "api.openai.com" || h.ends_with(".api.openai.com")) + .unwrap_or(false) +} + fn derive_base_path(url_path: &str) -> String { let stripped = url_path.trim_start_matches('/'); let normalized = stripped.trim_end_matches('/'); @@ -423,6 +430,16 @@ mod tests { assert_eq!(r, "https://opencode.ai/zen/go/v1/chat/completions"); } + #[test] + fn is_direct_openai_host_matches_only_openai() { + assert!(is_direct_openai_host("https://api.openai.com")); + assert!(is_direct_openai_host("https://api.openai.com/v1")); + assert!(is_direct_openai_host("https://eu.api.openai.com")); + assert!(!is_direct_openai_host("https://api.openai.com.local:8000")); + assert!(!is_direct_openai_host("https://localhost:1234")); + assert!(!is_direct_openai_host("https://router.huggingface.co/v1")); + } + #[test] fn derive_base_path_should_support_v1() { let r = derive_base_path("https://opencode.ai/zen/go/v1"); diff --git a/crates/goose/src/providers/openrouter.rs b/crates/goose/src/providers/openrouter.rs index 826e789a76fe..c3a6a481d433 100644 --- a/crates/goose/src/providers/openrouter.rs +++ b/crates/goose/src/providers/openrouter.rs @@ -40,7 +40,6 @@ pub const OPENROUTER_DOC_URL: &str = "https://openrouter.ai/models"; pub struct OpenRouterProvider { #[serde(skip)] api_client: ApiClient, - model: ModelConfig, supports_streaming: bool, #[serde(skip)] name: String, @@ -48,15 +47,8 @@ pub struct OpenRouterProvider { impl OpenRouterProvider { pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { - let model = crate::model_config::with_configured_fast_model( - model, - OPENROUTER_PROVIDER_NAME, - OPENROUTER_DEFAULT_FAST_MODEL, - )?; - let config = crate::config::Config::global(); let api_key: String = config.get_secret("OPENROUTER_API_KEY")?; let host: String = config @@ -70,7 +62,6 @@ impl OpenRouterProvider { Ok(Self { api_client, - model, supports_streaming: true, name: OPENROUTER_PROVIDER_NAME.to_string(), }) @@ -179,6 +170,7 @@ impl goose_providers::base::ProviderDescriptor for OpenRouterProvider { "Click 'Create' or use an existing API key", "Copy the key and paste it above", ]) + .with_fast_model(OPENROUTER_DEFAULT_FAST_MODEL) } } @@ -186,11 +178,10 @@ impl ProviderDef for OpenRouterProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -200,10 +191,6 @@ impl Provider for OpenRouterProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - /// Fetch supported models from OpenRouter API (only models with tool support) async fn fetch_supported_models(&self) -> Result, ProviderError> { let response = self @@ -252,12 +239,6 @@ impl Provider for OpenRouterProvider { Ok(models) } - async fn supports_cache_control(&self) -> bool { - self.model - .model_name - .starts_with(OPENROUTER_MODEL_PREFIX_ANTHROPIC) - } - async fn stream( &self, model_config: &ModelConfig, @@ -282,7 +263,7 @@ impl Provider for OpenRouterProvider { } } - if self.supports_cache_control().await { + if supports_cache_control(model_config) { payload = update_request_for_anthropic(&payload); } @@ -313,3 +294,9 @@ impl Provider for OpenRouterProvider { stream_openai_compat(response, log) } } + +fn supports_cache_control(model: &ModelConfig) -> bool { + model + .model_name + .starts_with(OPENROUTER_MODEL_PREFIX_ANTHROPIC) +} diff --git a/crates/goose/src/providers/pi_acp.rs b/crates/goose/src/providers/pi_acp.rs index 5432e845e545..ff249dcee751 100644 --- a/crates/goose/src/providers/pi_acp.rs +++ b/crates/goose/src/providers/pi_acp.rs @@ -11,7 +11,6 @@ use crate::config::{Config, GooseMode}; use crate::providers::base::{ current_working_dir, ProviderDef, ProviderDescriptor, ProviderMetadata, }; -use goose_providers::model::ModelConfig; pub(crate) const PI_ACP_PROVIDER_NAME: &str = "pi-acp"; const PI_ACP_DOC_URL: &str = "https://github.com/anthropics/pi"; @@ -44,15 +43,13 @@ impl ProviderDef for PiAcpProvider { type Provider = AcpProvider; fn from_env( - model: ModelConfig, extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Self::from_env_with_working_dir(model, extensions, current_working_dir(), tls_config) + Self::from_env_with_working_dir(extensions, current_working_dir(), tls_config) } fn from_env_with_working_dir( - model: ModelConfig, extensions: Vec, working_dir: PathBuf, _tls_config: Option, @@ -77,12 +74,14 @@ impl ProviderDef for PiAcpProvider { work_dir: working_dir, mcp_servers: extension_configs_to_mcp_servers(&extensions), session_mode_id: Some(mode_mapping[&goose_mode].clone()), + session_config_options: vec![], + model_config_option_id: None, mode_mapping, notification_callback: None, }; let metadata = Self::metadata(); - AcpProvider::connect(metadata.name, model, goose_mode, provider_config).await + AcpProvider::connect(metadata.name, goose_mode, provider_config).await }) } } diff --git a/crates/goose/src/providers/provider_registry.rs b/crates/goose/src/providers/provider_registry.rs index 9b5ca31749e4..8447a758b390 100644 --- a/crates/goose/src/providers/provider_registry.rs +++ b/crates/goose/src/providers/provider_registry.rs @@ -11,7 +11,6 @@ use std::sync::Arc; pub type ProviderConstructor = Arc< dyn Fn( - ModelConfig, Vec, Option, Option, @@ -55,7 +54,12 @@ impl ProviderEntry { (self.inventory_configured)() } - fn normalize_model_config(&self, mut model: ModelConfig) -> Result { + /// Apply provider-specific normalization to a model config: materialize + /// global defaults and backfill `context_limit` from the provider's known + /// models when the canonical registry didn't already resolve one. Used by + /// the agent/session layer to resolve effective limits (e.g. for custom + /// providers that declare explicit context limits in their config). + pub fn normalize_model_config(&self, mut model: ModelConfig) -> Result { model = crate::model_config::materialize_model_config(&self.metadata.name, model)?; if model.context_limit.is_none() { @@ -76,37 +80,19 @@ impl ProviderEntry { &self, extensions: Vec, ) -> Result> { - let model_config = crate::model_config::model_config_from_user_config( - &self.metadata.name, - &self.metadata.default_model, - )?; - let model_config = self.normalize_model_config(model_config)?; - (self.constructor)(model_config, extensions, None, self.tls_config.clone()).await + self.create(extensions).await } - pub async fn create( - &self, - model: ModelConfig, - extensions: Vec, - ) -> Result> { - let model = self.normalize_model_config(model)?; - (self.constructor)(model, extensions, None, self.tls_config.clone()).await + pub async fn create(&self, extensions: Vec) -> Result> { + (self.constructor)(extensions, None, self.tls_config.clone()).await } pub async fn create_with_working_dir( &self, - model: ModelConfig, extensions: Vec, working_dir: PathBuf, ) -> Result> { - let model = self.normalize_model_config(model)?; - (self.constructor)( - model, - extensions, - Some(working_dir), - self.tls_config.clone(), - ) - .await + (self.constructor)(extensions, Some(working_dir), self.tls_config.clone()).await } } @@ -147,19 +133,14 @@ impl ProviderRegistry { name, ProviderEntry { metadata, - constructor: Arc::new(|model, extensions, working_dir, tls_config| { + constructor: Arc::new(|extensions, working_dir, tls_config| { Box::pin(async move { let provider = match working_dir { Some(working_dir) => { - F::from_env_with_working_dir( - model, - extensions, - working_dir, - tls_config, - ) - .await? + F::from_env_with_working_dir(extensions, working_dir, tls_config) + .await? } - None => F::from_env(model, extensions, tls_config).await?, + None => F::from_env(extensions, tls_config).await?, }; Ok(Arc::new(provider) as Arc) }) @@ -187,7 +168,7 @@ impl ProviderRegistry { inventory_identity: G, ) where P: ProviderDef + 'static, - F: Fn(ModelConfig, Option) -> Result + Send + Sync + 'static, + F: Fn(Option) -> Result + Send + Sync + 'static, G: Fn() -> Result + Send + Sync + 'static, { self.register_with_name_impl::( @@ -210,7 +191,7 @@ impl ProviderRegistry { inventory_configured: H, ) where P: ProviderDef + 'static, - F: Fn(ModelConfig, Option) -> Result + Send + Sync + 'static, + F: Fn(Option) -> Result + Send + Sync + 'static, G: Fn() -> Result + Send + Sync + 'static, H: Fn() -> bool + Send + Sync + 'static, { @@ -234,7 +215,7 @@ impl ProviderRegistry { inventory_configured: Option, ) where P: ProviderDef + 'static, - F: Fn(ModelConfig, Option) -> Result + Send + Sync + 'static, + F: Fn(Option) -> Result + Send + Sync + 'static, G: Fn() -> Result + Send + Sync + 'static, { let base_metadata = P::metadata(); @@ -316,6 +297,7 @@ impl ProviderRegistry { config_keys, setup_steps: config.setup_steps.clone(), model_selection_hint: None, + fast_model: config.fast_model.clone(), }; let inventory_config_keys = custom_metadata.config_keys.clone(); let default_inventory_configured = Arc::new(move || { @@ -329,8 +311,8 @@ impl ProviderRegistry { config.name.clone(), ProviderEntry { metadata: custom_metadata, - constructor: Arc::new(move |model, _extensions, _working_dir, tls_config| { - let result = constructor(model, tls_config); + constructor: Arc::new(move |_extensions, _working_dir, tls_config| { + let result = constructor(tls_config); Box::pin(async move { let provider = result?; Ok(Arc::new(provider) as Arc) @@ -363,7 +345,6 @@ impl ProviderRegistry { pub async fn create( &self, name: &str, - model: ModelConfig, extensions: Vec, ) -> Result> { let entry = self @@ -371,7 +352,7 @@ impl ProviderRegistry { .get(name) .ok_or_else(|| anyhow::anyhow!("Unknown provider: {}", name))?; - entry.create(model, extensions).await + entry.create(extensions).await } pub fn all_metadata_with_types(&self) -> Vec<(ProviderMetadata, ProviderType)> { @@ -424,7 +405,7 @@ mod tests { &test_config(), ProviderType::Declarative, false, - |_, _| unreachable!("constructor is not used by this test"), + |_| unreachable!("constructor is not used by this test"), || Ok(InventoryIdentityInput::new("custom_hf", "huggingface")), || false, ); diff --git a/crates/goose/src/providers/provider_test.rs b/crates/goose/src/providers/provider_test.rs index 7bb979a76279..5284f2a5d3e5 100644 --- a/crates/goose/src/providers/provider_test.rs +++ b/crates/goose/src/providers/provider_test.rs @@ -15,7 +15,7 @@ pub async fn test_provider_configuration( .with_toolshim(toolshim_enabled) .with_toolshim_model(toolshim_model); - let provider = create(provider_name, model_config, Vec::new()).await?; + let provider = create(provider_name, Vec::new()).await?; let messages = vec![Message::user().with_text("What is the weather like in San Francisco today?")]; @@ -26,10 +26,9 @@ pub async fn test_provider_configuration( vec![] }; - let provider_model_config = provider.get_model_config(); let mut stream = provider .stream( - &provider_model_config, + &model_config, "test-session-id", "You are an AI agent called goose. You use tools of connected extensions to solve problems.", &messages, diff --git a/crates/goose/src/providers/sagemaker_tgi.rs b/crates/goose/src/providers/sagemaker_tgi.rs index cc994b6cf319..18108bc26e25 100644 --- a/crates/goose/src/providers/sagemaker_tgi.rs +++ b/crates/goose/src/providers/sagemaker_tgi.rs @@ -34,14 +34,12 @@ pub struct SageMakerTgiProvider { #[serde(skip)] sagemaker_client: SageMakerClient, endpoint_name: String, - model: ModelConfig, #[serde(skip)] name: String, } impl SageMakerTgiProvider { pub async fn from_env( - model: ModelConfig, _tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -91,12 +89,16 @@ impl SageMakerTgiProvider { Ok(Self { sagemaker_client, endpoint_name, - model, name: SAGEMAKER_TGI_PROVIDER_NAME.to_string(), }) } - fn create_tgi_request(&self, system: &str, messages: &[Message]) -> Result { + fn create_tgi_request( + &self, + model: &ModelConfig, + system: &str, + messages: &[Message], + ) -> Result { // Create a simplified prompt for TGI models using recent user and assistant messages. // Uses a minimal system prompt and avoids HTML or tool-related formatting. let mut prompt = String::new(); @@ -155,8 +157,8 @@ impl SageMakerTgiProvider { let request = json!({ "inputs": prompt, "parameters": { - "max_new_tokens": self.model.max_output_tokens(), - "temperature": self.model.temperature.unwrap_or(0.7), + "max_new_tokens": model.max_output_tokens(), + "temperature": model.temperature.unwrap_or(0.7), "do_sample": true, "return_full_text": false } @@ -300,11 +302,10 @@ impl ProviderDef for SageMakerTgiProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -314,10 +315,6 @@ impl Provider for SageMakerTgiProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn stream( &self, model_config: &ModelConfig, @@ -333,9 +330,11 @@ impl Provider for SageMakerTgiProvider { }; let model_name = &model_config.model_name; - let request_payload = self.create_tgi_request(system, messages).map_err(|e| { - ProviderError::RequestFailed(format!("Failed to create request: {}", e)) - })?; + let request_payload = self + .create_tgi_request(model_config, system, messages) + .map_err(|e| { + ProviderError::RequestFailed(format!("Failed to create request: {}", e)) + })?; let response = self .with_retry(|| self.invoke_endpoint(session_id, request_payload.clone())) @@ -356,7 +355,7 @@ impl Provider for SageMakerTgiProvider { "messages": messages, "tools": tools }); - let mut log = start_log(&self.model, &debug_payload)?; + let mut log = start_log(model_config, &debug_payload)?; log.write( &serde_json::to_value(&message).unwrap_or_default(), Some(&usage), diff --git a/crates/goose/src/providers/snowflake.rs b/crates/goose/src/providers/snowflake.rs index be7fba93ce39..0bb5c174043b 100644 --- a/crates/goose/src/providers/snowflake.rs +++ b/crates/goose/src/providers/snowflake.rs @@ -52,7 +52,6 @@ impl SnowflakeAuth { pub struct SnowflakeProvider { #[serde(skip)] api_client: ApiClient, - model: ModelConfig, image_format: ImageFormat, #[serde(skip)] name: String, @@ -60,7 +59,6 @@ pub struct SnowflakeProvider { impl SnowflakeProvider { pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -111,7 +109,6 @@ impl SnowflakeProvider { Ok(Self { api_client, - model, image_format: ImageFormat::OpenAi, name: SNOWFLAKE_PROVIDER_NAME.to_string(), }) @@ -324,11 +321,10 @@ impl ProviderDef for SnowflakeProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -338,10 +334,6 @@ impl Provider for SnowflakeProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn fetch_supported_models(&self) -> Result, ProviderError> { Ok(SNOWFLAKE_KNOWN_MODELS .iter() @@ -364,7 +356,7 @@ impl Provider for SnowflakeProvider { }; let payload = create_request(model_config, system, messages, tools)?; - let mut log = start_log(&self.model, &payload)?; + let mut log = start_log(model_config, &payload)?; let response = self .with_retry(|| async { diff --git a/crates/goose/src/providers/testprovider.rs b/crates/goose/src/providers/testprovider.rs index a1298812759e..6a6620edbc52 100644 --- a/crates/goose/src/providers/testprovider.rs +++ b/crates/goose/src/providers/testprovider.rs @@ -156,7 +156,6 @@ impl ProviderDef for TestProvider { type Provider = Self; fn from_env( - _model: ModelConfig, _extensions: Vec, _tls_config: Option, ) -> BoxFuture<'static, Result> { @@ -219,10 +218,6 @@ impl Provider for TestProvider { } } } - - fn get_model_config(&self) -> ModelConfig { - ModelConfig::new_or_fail("test-model") - } } #[cfg(test)] @@ -236,7 +231,6 @@ mod tests { #[derive(Clone)] struct MockProvider { - model_config: ModelConfig, response: String, } @@ -268,10 +262,6 @@ mod tests { let usage = ProviderUsage::new("mock-model".to_string(), Usage::default()); Ok(stream_from_single_message(message, usage)) } - - fn get_model_config(&self) -> ModelConfig { - self.model_config.clone() - } } #[tokio::test] @@ -283,13 +273,12 @@ mod tests { ); let mock = Arc::new(MockProvider { - model_config: ModelConfig::new_or_fail("mock-model"), response: "Hello, world!".to_string(), }); { let test_provider = TestProvider::new_recording(mock, &temp_file); - let model_config = test_provider.get_model_config(); + let model_config = ModelConfig::new("test-model"); let result = test_provider .complete( @@ -314,7 +303,7 @@ mod tests { { let replay_provider = TestProvider::new_replaying(&temp_file).unwrap(); - let model_config = replay_provider.get_model_config(); + let model_config = ModelConfig::new("test-model"); let result = replay_provider .complete( @@ -346,7 +335,7 @@ mod tests { ); let replay_provider = TestProvider::new_replaying(&temp_file).unwrap(); - let model_config = replay_provider.get_model_config(); + let model_config = ModelConfig::new("test-model"); let result = replay_provider .complete( diff --git a/crates/goose/src/providers/tetrate.rs b/crates/goose/src/providers/tetrate.rs index 3d906c10db89..38ac05805069 100644 --- a/crates/goose/src/providers/tetrate.rs +++ b/crates/goose/src/providers/tetrate.rs @@ -40,7 +40,6 @@ pub const TETRATE_KNOWN_MODELS: &[&str] = &[ pub struct TetrateProvider { #[serde(skip)] api_client: ApiClient, - model: ModelConfig, supports_streaming: bool, #[serde(skip)] name: String, @@ -48,7 +47,6 @@ pub struct TetrateProvider { impl TetrateProvider { pub async fn from_env( - model: ModelConfig, tls_config: Option, ) -> Result { let config = crate::config::Config::global(); @@ -64,7 +62,6 @@ impl TetrateProvider { Ok(Self { api_client, - model, supports_streaming: true, name: TETRATE_PROVIDER_NAME.to_string(), }) @@ -119,11 +116,10 @@ impl ProviderDef for TetrateProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { - Box::pin(Self::from_env(model, tls_config)) + Box::pin(Self::from_env(tls_config)) } } @@ -133,10 +129,6 @@ impl Provider for TetrateProvider { &self.name } - fn get_model_config(&self) -> ModelConfig { - self.model.clone() - } - async fn stream( &self, model_config: &ModelConfig, diff --git a/crates/goose/src/providers/toolshim.rs b/crates/goose/src/providers/toolshim.rs index 4e2b87570edc..d16c51c53ff9 100644 --- a/crates/goose/src/providers/toolshim.rs +++ b/crates/goose/src/providers/toolshim.rs @@ -571,7 +571,7 @@ impl LocalInterpreter { .with_toolshim(false) .with_toolshim_model(None); - let provider = crate::providers::init::create("local", model_config, vec![]) + let provider = crate::providers::init::create("local", vec![]) .await .map_err(|e| { ProviderError::RequestFailed(format!( @@ -581,13 +581,7 @@ impl LocalInterpreter { let request_messages = vec![Message::user().with_text(format_instruction)]; let mut stream = provider - .stream( - &provider.get_model_config(), - "toolshim-local", - "", - &request_messages, - &[], - ) + .stream(&model_config, "toolshim-local", "", &request_messages, &[]) .await?; let mut content = String::new(); diff --git a/crates/goose/src/providers/xai.rs b/crates/goose/src/providers/xai.rs index 74a91fbc5754..7488e713d353 100644 --- a/crates/goose/src/providers/xai.rs +++ b/crates/goose/src/providers/xai.rs @@ -3,7 +3,6 @@ use super::base::{ConfigKey, ProviderDef, ProviderMetadata}; use super::openai_compatible::OpenAiCompatibleProvider; use anyhow::Result; use futures::future::BoxFuture; -use goose_providers::model::ModelConfig; const XAI_PROVIDER_NAME: &str = "xai"; pub const XAI_API_HOST: &str = "https://api.x.ai/v1"; @@ -54,7 +53,6 @@ impl ProviderDef for XaiProvider { type Provider = OpenAiCompatibleProvider; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { @@ -71,7 +69,6 @@ impl ProviderDef for XaiProvider { Ok(OpenAiCompatibleProvider::new( XAI_PROVIDER_NAME.to_string(), api_client, - model, String::new(), )) }) diff --git a/crates/goose/src/providers/xai_oauth.rs b/crates/goose/src/providers/xai_oauth.rs index 19b44963081a..8b675bd0e3ac 100644 --- a/crates/goose/src/providers/xai_oauth.rs +++ b/crates/goose/src/providers/xai_oauth.rs @@ -703,10 +703,6 @@ impl Provider for XaiOAuthProvider { self.inner.get_name() } - fn get_model_config(&self) -> ModelConfig { - self.inner.get_model_config() - } - async fn stream( &self, model_config: &ModelConfig, @@ -781,7 +777,6 @@ impl ProviderDef for XaiOAuthProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, tls_config: Option, ) -> BoxFuture<'static, Result> { @@ -802,7 +797,6 @@ impl ProviderDef for XaiOAuthProvider { let inner = OpenAiCompatibleProvider::new( XAI_OAUTH_PROVIDER_NAME.to_string(), api_client, - model, String::new(), ); diff --git a/crates/goose/src/scheduler.rs b/crates/goose/src/scheduler.rs index 6c42cbd0db38..8466893ca021 100644 --- a/crates/goose/src/scheduler.rs +++ b/crates/goose/src/scheduler.rs @@ -858,8 +858,10 @@ async fn execute_job( agent.add_extension(ext.clone(), &session.id).await?; } - let agent_provider = create(&provider_name, model_config, extensions).await?; - agent.update_provider(agent_provider, &session.id).await?; + let agent_provider = create(&provider_name, extensions).await?; + agent + .update_provider(agent_provider, model_config, &session.id) + .await?; let mut jobs_guard = jobs.lock().await; if let Some((_, job_def)) = jobs_guard.get_mut(job_id.as_str()) { diff --git a/crates/goose/src/security/adversary_inspector.rs b/crates/goose/src/security/adversary_inspector.rs index e7c6d320d3fc..d0c038ca590e 100644 --- a/crates/goose/src/security/adversary_inspector.rs +++ b/crates/goose/src/security/adversary_inspector.rs @@ -1,7 +1,7 @@ use anyhow::Result; use async_trait::async_trait; use chrono::Utc; -use std::sync::OnceLock; +use std::sync::{Arc, OnceLock}; use crate::agents::types::SharedProvider; use crate::config::paths::Paths; @@ -13,6 +13,28 @@ use crate::utils::safe_truncate; const DEFAULT_TOOLS: &[&str] = &["shell", "computercontroller__automation_script"]; +async fn resolve_model_config( + session_manager: &crate::session::SessionManager, + session_id: &str, +) -> Result { + if !session_id.is_empty() { + if let Ok(session) = session_manager.get_session(session_id, false).await { + if let Some(model_config) = session.model_config { + return Ok(model_config); + } + } + } + + let config = crate::config::Config::global(); + let provider_name = config + .get_goose_provider() + .map_err(|_| anyhow::anyhow!("missing provider"))?; + let model_name = config + .get_goose_model() + .map_err(|_| anyhow::anyhow!("missing model"))?; + crate::model_config::model_config_from_user_config(&provider_name, &model_name) +} + const DEFAULT_RULES: &str = r#"BLOCK if the command: - Exfiltrates data (curl/wget posting to unknown URLs, piping secrets out) - Is destructive beyond the project scope (rm -rf /, modifying system files) @@ -50,22 +72,32 @@ struct AdversaryConfig { /// If the review fails, the inspector fails open (allows the tool call). pub struct AdversaryInspector { provider: SharedProvider, + session_manager: Arc, config: OnceLock>, config_path: Option, } impl AdversaryInspector { - pub fn new(provider: SharedProvider) -> Self { + pub fn new( + provider: SharedProvider, + session_manager: Arc, + ) -> Self { Self { provider, + session_manager, config: OnceLock::new(), config_path: None, } } - pub fn with_config_dir(provider: SharedProvider, config_dir: std::path::PathBuf) -> Self { + pub fn with_config_dir( + provider: SharedProvider, + session_manager: Arc, + config_dir: std::path::PathBuf, + ) -> Self { Self { provider, + session_manager, config: OnceLock::new(), config_path: Some(config_dir.join("adversary.md")), } @@ -240,6 +272,7 @@ impl AdversaryInspector { async fn consult_llm( &self, + session_id: &str, tool_description: &str, original_task: &str, recent_messages: &[String], @@ -290,11 +323,13 @@ impl AdversaryInspector { )]; let conversation = Conversation::new_unvalidated(check_messages); - let model_config = provider.get_model_config(); + let model_config = resolve_model_config(&self.session_manager, session_id) + .await + .map_err(|e| anyhow::anyhow!("Could not resolve model config: {}", e))?; let (response, _usage) = provider .complete( &model_config, - "", + session_id, system_prompt, conversation.messages(), &[], @@ -358,7 +393,7 @@ impl ToolInspector for AdversaryInspector { async fn inspect( &self, - _session_id: &str, + session_id: &str, tool_requests: &[ToolRequest], messages: &[Message], _goose_mode: GooseMode, @@ -392,6 +427,7 @@ impl ToolInspector for AdversaryInspector { match self .consult_llm( + session_id, &tool_description, &original_task, &recent_messages, @@ -625,7 +661,14 @@ mod tests { let tmp = tempfile::tempdir().unwrap(); let provider: SharedProvider = Arc::new(Mutex::new(None)); - let inspector = AdversaryInspector::with_config_dir(provider, tmp.path().to_path_buf()); + let session_manager = Arc::new(crate::session::SessionManager::new( + tmp.path().to_path_buf(), + )); + let inspector = AdversaryInspector::with_config_dir( + provider, + session_manager, + tmp.path().to_path_buf(), + ); assert!(!inspector.is_enabled()); let request = ToolRequest { diff --git a/crates/goose/src/session/session_manager.rs b/crates/goose/src/session/session_manager.rs index 1abdf14976b3..6b7971211657 100644 --- a/crates/goose/src/session/session_manager.rs +++ b/crates/goose/src/session/session_manager.rs @@ -495,6 +495,21 @@ impl SessionManager { return Ok(Some(self.system_generated_name_update(id, name).await?)); } + let model_config = match session.model_config.clone() { + Some(model_config) => model_config, + None => { + let model_name = + crate::config::Config::global() + .get_goose_model() + .map_err(|_| { + anyhow::anyhow!("Could not resolve model config: missing model") + })?; + crate::model_config::model_config_from_user_config( + provider.get_name(), + &model_name, + )? + } + }; let conversation = session .conversation .ok_or_else(|| anyhow::anyhow!("No messages found"))?; @@ -506,7 +521,8 @@ impl SessionManager { .count(); if user_message_count <= MSG_COUNT_FOR_SESSION_NAME_GENERATION { - let name = generate_session_name(provider.as_ref(), id, &conversation).await?; + let name = + generate_session_name(provider.as_ref(), &model_config, id, &conversation).await?; return Ok(Some(self.system_generated_name_update(id, name).await?)); } Ok(None) @@ -2089,9 +2105,7 @@ mod tests { const NUM_CONCURRENT_SESSIONS: i32 = 10; const GENERATED_SESSION_NAME: &str = "Generated session name"; - struct NamingTestProvider { - model_config: ModelConfig, - } + struct NamingTestProvider; #[async_trait::async_trait] impl Provider for NamingTestProvider { @@ -2107,15 +2121,12 @@ mod tests { _messages: &[Message], _tools: &[rmcp::model::Tool], ) -> std::result::Result { - unimplemented!("session naming calls complete_fast") + unimplemented!("session naming calls complete") } - fn get_model_config(&self) -> ModelConfig { - self.model_config.clone() - } - - async fn complete_fast( + async fn complete( &self, + _model_config: &ModelConfig, _session_id: &str, _system: &str, _messages: &[Message], @@ -2129,9 +2140,7 @@ mod tests { } fn naming_test_provider() -> Arc { - Arc::new(NamingTestProvider { - model_config: ModelConfig::new("test-model").unwrap(), - }) + Arc::new(NamingTestProvider) } fn test_recipe(title: &str) -> Recipe { diff --git a/crates/goose/src/session/session_naming.rs b/crates/goose/src/session/session_naming.rs index 74321016aaf9..8391d252c21e 100644 --- a/crates/goose/src/session/session_naming.rs +++ b/crates/goose/src/session/session_naming.rs @@ -112,6 +112,7 @@ fn get_preprompt_context(messages: &Conversation) -> String { /// Creates a prompt asking for a concise description in 4 words or less. pub(crate) async fn generate_session_name( provider: &dyn Provider, + model_config: &goose_providers::model::ModelConfig, session_id: &str, messages: &Conversation, ) -> Result { @@ -144,9 +145,15 @@ pub(crate) async fn generate_session_name( SESSION_NAME_SUFFIX, ); let message = Message::user().with_text(&user_text); - let result = provider - .complete_fast(session_id, &system, &[message], &[]) - .await?; + let result = crate::model_config::complete_fast( + provider, + model_config, + session_id, + &system, + &[message], + &[], + ) + .await?; let raw: String = result .0 diff --git a/crates/goose/tests/acp_common_tests/mod.rs b/crates/goose/tests/acp_common_tests/mod.rs index 49464e8278b5..0065a7bd10d9 100644 --- a/crates/goose/tests/acp_common_tests/mod.rs +++ b/crates/goose/tests/acp_common_tests/mod.rs @@ -429,7 +429,7 @@ pub async fn run_fs_write_text_file_true() { pub async fn run_initialize_doesnt_hit_provider() { let provider_factory: AcpProviderFactory = - Arc::new(|_, _, _, _| Box::pin(async { Err(anyhow::anyhow!("no provider configured")) })); + Arc::new(|_, _, _| Box::pin(async { Err(anyhow::anyhow!("no provider configured")) })); let openai = OpenAiFixture::new(vec![], C::expected_session_id()).await; let config = TestConnectionConfig { diff --git a/crates/goose/tests/acp_custom_requests_test.rs b/crates/goose/tests/acp_custom_requests_test.rs index 3c9c87c859a9..c69ea6dd0425 100644 --- a/crates/goose/tests/acp_custom_requests_test.rs +++ b/crates/goose/tests/acp_custom_requests_test.rs @@ -47,7 +47,6 @@ fn write_acp_global_config(contents: &str) -> PathBuf { struct MockProvider { name: String, - model_config: ModelConfig, recommended_models: Vec, supported_models: Vec, } @@ -69,11 +68,10 @@ impl Provider for MockProvider { unimplemented!() } - fn get_model_config(&self) -> ModelConfig { - self.model_config.clone() - } - - async fn fetch_recommended_models(&self) -> Result, ProviderError> { + async fn fetch_recommended_models( + &self, + _toolshim: bool, + ) -> Result, ProviderError> { Ok(self.recommended_models.clone()) } @@ -1040,11 +1038,10 @@ fn test_custom_provider_supported_models_lists_raw_provider_models() { run_test(async move { let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await; let provider_factory: AcpProviderFactory = - Arc::new(|provider_name, model_config, _extensions, _working_dir| { + Arc::new(|provider_name, _extensions, _working_dir| { Box::pin(async move { Ok(Arc::new(MockProvider { name: provider_name, - model_config, recommended_models: vec!["canonical-filtered-model".to_string()], supported_models: vec![ "goose-claude-opus-4-8".to_string(), diff --git a/crates/goose/tests/acp_fixtures/mod.rs b/crates/goose/tests/acp_fixtures/mod.rs index cde43e1a8b0e..c799dcf52348 100644 --- a/crates/goose/tests/acp_fixtures/mod.rs +++ b/crates/goose/tests/acp_fixtures/mod.rs @@ -347,22 +347,19 @@ pub async fn spawn_acp_server_in_process( write_global_test_config(&config_path, openai_base_url); let provider_factory = provider_factory.unwrap_or_else(|| { let base_url = openai_base_url.to_string(); - Arc::new( - move |_provider_name, model_config, _extensions, _working_dir| { - let base_url = base_url.clone(); - Box::pin(async move { - let api_client = ApiClient::new_with_tls( - base_url, - ApiAuthMethod::BearerToken("test-key".to_string()), - None, - ) - .unwrap(); - let provider: Arc = - Arc::new(OpenAiProvider::new(api_client, model_config)); - Ok(provider) - }) - }, - ) + Arc::new(move |_provider_name, _extensions, _working_dir| { + let base_url = base_url.clone(); + Box::pin(async move { + let api_client = ApiClient::new_with_tls( + base_url, + ApiAuthMethod::BearerToken("test-key".to_string()), + None, + ) + .unwrap(); + let provider: Arc = Arc::new(OpenAiProvider::new(api_client)); + Ok(provider) + }) + }) }); let agent = GooseAcpAgent::new(GooseAcpAgentOptions { diff --git a/crates/goose/tests/acp_fixtures/provider.rs b/crates/goose/tests/acp_fixtures/provider.rs index c11dc9247fcd..3ce4ca3025bd 100644 --- a/crates/goose/tests/acp_fixtures/provider.rs +++ b/crates/goose/tests/acp_fixtures/provider.rs @@ -76,7 +76,7 @@ impl AcpProviderSession { .unwrap() .get(session_id.as_ref()) .cloned() - .unwrap_or_else(|| provider.get_model_config()); + .unwrap_or_else(|| ModelConfig::new(TEST_MODEL)); let mut stream = provider .stream(&model_config, &session_id, "", &[message], &[]) .await?; @@ -183,6 +183,8 @@ impl Connection for AcpProviderConnection { work_dir: cwd_path.clone(), mcp_servers, session_mode_id: None, + session_config_options: vec![], + model_config_option_id: None, mode_mapping: GooseMode::VARIANTS .iter() .map(|v| { @@ -203,7 +205,6 @@ impl Connection for AcpProviderConnection { }; let provider = AcpProvider::connect_with_transport( "acp-test".to_string(), - ModelConfig::new(TEST_MODEL).unwrap(), goose_mode, provider_config, transport, @@ -237,7 +238,7 @@ impl Connection for AcpProviderConnection { let provider = provider.as_ref().unwrap(); let available_models = provider.fetch_supported_models().await?; Some(SessionModelState::new( - ModelId::new(provider.get_model_config().model_name.clone()), + ModelId::new(TEST_MODEL.to_string()), available_models .into_iter() .map(|model_id| ModelInfo::new(ModelId::new(model_id.clone()), model_id)) diff --git a/crates/goose/tests/acp_secret_cache_invalidation_test.rs b/crates/goose/tests/acp_secret_cache_invalidation_test.rs index bce30e4a6473..9f4bf88f9eb1 100644 --- a/crates/goose/tests/acp_secret_cache_invalidation_test.rs +++ b/crates/goose/tests/acp_secret_cache_invalidation_test.rs @@ -17,7 +17,6 @@ use std::sync::Arc; struct MockProvider { name: String, - model_config: ModelConfig, } #[async_trait::async_trait] @@ -37,21 +36,19 @@ impl Provider for MockProvider { unimplemented!() } - fn get_model_config(&self) -> ModelConfig { - self.model_config.clone() - } - - async fn fetch_recommended_models(&self) -> Result, ProviderError> { + async fn fetch_recommended_models( + &self, + _toolshim: bool, + ) -> Result, ProviderError> { Ok(vec!["claude-3-5-haiku-latest".to_string()]) } } fn mock_provider_factory() -> goose::acp::server::AcpProviderFactory { - Arc::new(|provider_name, model_config, _extensions, _working_dir| { + Arc::new(|provider_name, _extensions, _working_dir| { Box::pin(async move { Ok(Arc::new(MockProvider { name: provider_name, - model_config, }) as Arc) }) }) diff --git a/crates/goose/tests/adversary_inspector_tests.rs b/crates/goose/tests/adversary_inspector_tests.rs index cd11bd563ad9..8ace7da80c95 100644 --- a/crates/goose/tests/adversary_inspector_tests.rs +++ b/crates/goose/tests/adversary_inspector_tests.rs @@ -30,7 +30,13 @@ async fn test_adversary_disabled_without_config_file() { let tmp = tempfile::tempdir().unwrap(); let provider = Arc::new(Mutex::new(None)); - let inspector = AdversaryInspector::with_config_dir(provider, tmp.path().to_path_buf()); + let inspector = AdversaryInspector::with_config_dir( + provider, + Arc::new(goose::session::SessionManager::new( + tmp.path().to_path_buf(), + )), + tmp.path().to_path_buf(), + ); assert_eq!(inspector.name(), "adversary"); assert!(!inspector.is_enabled()); @@ -58,7 +64,13 @@ async fn test_adversary_enabled_default_tools() { write_adversary_md(tmp.path(), "BLOCK everything for testing"); let provider = Arc::new(Mutex::new(None)); - let inspector = AdversaryInspector::with_config_dir(provider, tmp.path().to_path_buf()); + let inspector = AdversaryInspector::with_config_dir( + provider, + Arc::new(goose::session::SessionManager::new( + tmp.path().to_path_buf(), + )), + tmp.path().to_path_buf(), + ); assert!(inspector.is_enabled()); @@ -116,7 +128,13 @@ async fn test_adversary_custom_tool_filter() { ); let provider = Arc::new(Mutex::new(None)); - let inspector = AdversaryInspector::with_config_dir(provider, tmp.path().to_path_buf()); + let inspector = AdversaryInspector::with_config_dir( + provider, + Arc::new(goose::session::SessionManager::new( + tmp.path().to_path_buf(), + )), + tmp.path().to_path_buf(), + ); assert!(inspector.is_enabled()); diff --git a/crates/goose/tests/agent.rs b/crates/goose/tests/agent.rs index c0be51adc2fd..3a932e073218 100644 --- a/crates/goose/tests/agent.rs +++ b/crates/goose/tests/agent.rs @@ -373,6 +373,7 @@ mod tests { config_keys: vec![], setup_steps: vec![], model_selection_hint: None, + fast_model: None, } } } @@ -381,7 +382,6 @@ mod tests { type Provider = Self; fn from_env( - _model: ModelConfig, _extensions: Vec, _tls_config: Option, ) -> futures::future::BoxFuture<'static, anyhow::Result> { @@ -411,10 +411,6 @@ mod tests { Ok(stream_from_single_message(message, usage)) } - fn get_model_config(&self) -> ModelConfig { - ModelConfig::new("mock-model").unwrap() - } - fn get_name(&self) -> &str { "mock-test" } @@ -437,7 +433,9 @@ mod tests { ) .await?; - agent.update_provider(provider, &session.id).await?; + agent + .update_provider(provider, ModelConfig::new("mock-model"), &session.id) + .await?; let session_config = SessionConfig { id: session.id, @@ -546,6 +544,7 @@ mod tests { config_keys: vec![], setup_steps: vec![], model_selection_hint: None, + fast_model: None, } } } @@ -554,7 +553,6 @@ mod tests { type Provider = Self; fn from_env( - _model: ModelConfig, _extensions: Vec, _tls_config: Option, ) -> futures::future::BoxFuture<'static, anyhow::Result> { @@ -588,10 +586,6 @@ mod tests { Ok(stream_from_single_message(message, usage)) } - fn get_model_config(&self) -> ModelConfig { - ModelConfig::new("mock-model").unwrap() - } - fn get_name(&self) -> &str { "mock-summarization" } @@ -624,7 +618,9 @@ mod tests { ) .await?; - agent.update_provider(provider, &session.id).await?; + agent + .update_provider(provider, ModelConfig::new("mock-model"), &session.id) + .await?; // Pre-populate 13 tool pairs (need > cutoff + batch_size = 12 to trigger). // Timestamps in the past so DB ordering places summaries before current turn. @@ -901,6 +897,7 @@ mod tests { config_keys: vec![], setup_steps: vec![], model_selection_hint: None, + fast_model: None, } } } @@ -909,7 +906,6 @@ mod tests { type Provider = Self; fn from_env( - _model: ModelConfig, _extensions: Vec, _tls_config: Option, ) -> futures::future::BoxFuture<'static, anyhow::Result> { @@ -978,10 +974,6 @@ mod tests { } } - fn get_model_config(&self) -> ModelConfig { - ModelConfig::new("mock-model").unwrap() - } - fn get_name(&self) -> &str { "multi-step-mock" } @@ -1013,7 +1005,9 @@ mod tests { .await?; let session_id = session.id.clone(); - agent.update_provider(provider, &session_id).await?; + agent + .update_provider(provider, ModelConfig::new("mock-model"), &session_id) + .await?; // ── Single reply: tool call (call 0) → text stream (call 1) → cancelled text (call 2) // max_turns=3 allows all three provider calls within one reply(). @@ -1174,6 +1168,7 @@ mod tests { config_keys: vec![], setup_steps: vec![], model_selection_hint: None, + fast_model: None, } } } @@ -1182,7 +1177,6 @@ mod tests { type Provider = Self; fn from_env( - _model: ModelConfig, _extensions: Vec, _tls_config: Option, ) -> futures::future::BoxFuture<'static, anyhow::Result> { @@ -1210,10 +1204,6 @@ mod tests { Ok(stream_from_single_message(message, usage)) } - fn get_model_config(&self) -> ModelConfig { - ModelConfig::new("mock-model").unwrap() - } - fn get_name(&self) -> &str { "goal-mock" } @@ -1249,7 +1239,13 @@ mod tests { ) .await?; - agent.update_provider(provider.clone(), &session.id).await?; + agent + .update_provider( + provider.clone(), + ModelConfig::new("mock-model"), + &session.id, + ) + .await?; agent .set_goal(Some("Ensure the sky is blue".to_string())) .await; @@ -1325,7 +1321,13 @@ mod tests { ) .await?; - agent.update_provider(provider.clone(), &session.id).await?; + agent + .update_provider( + provider.clone(), + ModelConfig::new("mock-model"), + &session.id, + ) + .await?; let session_config = SessionConfig { id: session.id.clone(), @@ -1415,7 +1417,13 @@ mod tests { GooseMode::default(), ) .await?; - agent.update_provider(provider.clone(), &session.id).await?; + agent + .update_provider( + provider.clone(), + ModelConfig::new("mock-model"), + &session.id, + ) + .await?; let session_config = SessionConfig { id: session.id.clone(), @@ -1473,7 +1481,13 @@ mod tests { GooseMode::default(), ) .await?; - agent.update_provider(provider.clone(), &session.id).await?; + agent + .update_provider( + provider.clone(), + ModelConfig::new("mock-model"), + &session.id, + ) + .await?; let session_config = SessionConfig { id: session.id.clone(), @@ -1545,10 +1559,6 @@ mod tests { Ok(stream_from_single_message(message, usage)) } - fn get_model_config(&self) -> ModelConfig { - ModelConfig::new("mock-model").unwrap() - } - fn get_name(&self) -> &str { "fixed-usage-mock" } @@ -1599,7 +1609,13 @@ mod tests { .await?; let session_id = session.id.clone(); - agent.update_provider(provider.clone(), &session_id).await?; + agent + .update_provider( + provider.clone(), + ModelConfig::new("mock-model"), + &session_id, + ) + .await?; run_turn(&agent, &session_id, "Turn 1").await?; let after_1 = session_manager.get_session(&session_id, false).await?; diff --git a/crates/goose/tests/compaction.rs b/crates/goose/tests/compaction.rs index 3b95b2d11fe6..5d4d0769f72e 100644 --- a/crates/goose/tests/compaction.rs +++ b/crates/goose/tests/compaction.rs @@ -170,10 +170,6 @@ impl Provider for MockCompactionProvider { Ok(stream_from_single_message(message, usage)) } - fn get_model_config(&self) -> ModelConfig { - ModelConfig::new("mock-model").unwrap() - } - fn get_name(&self) -> &str { "mock-compaction" } @@ -191,6 +187,7 @@ impl goose::providers::base::ProviderDescriptor for MockCompactionProvider { config_keys: vec![], setup_steps: vec![], model_selection_hint: None, + fast_model: None, } } } @@ -199,7 +196,6 @@ impl ProviderDef for MockCompactionProvider { type Provider = Self; fn from_env( - _model: ModelConfig, _extensions: Vec, _tls_config: Option, ) -> futures::future::BoxFuture<'static, anyhow::Result> { @@ -348,7 +344,9 @@ async fn test_manual_compaction_updates_token_counts_and_conversation() -> Resul // Setup mock provider let provider = Arc::new(MockCompactionProvider::new()); - agent.update_provider(provider, &session.id).await?; + agent + .update_provider(provider, ModelConfig::new("mock-model"), &session.id) + .await?; // Execute manual compaction let result = agent.execute_command("/compact", &session.id).await?; @@ -444,7 +442,9 @@ async fn test_auto_compaction_during_reply() -> Result<()> { // Setup mock provider (no context limit enforcement) let provider = Arc::new(MockCompactionProvider::new()); - agent.update_provider(provider, &session.id).await?; + agent + .update_provider(provider, ModelConfig::new("mock-model"), &session.id) + .await?; // Trigger a reply // Expected tokens for reply: @@ -600,7 +600,9 @@ async fn test_context_limit_recovery_compaction() -> Result<()> { // Setup mock provider with context limit of 20000 tokens // Initial context (6000 system + 15400 messages = 21400) exceeds this limit let provider = Arc::new(MockCompactionProvider::new()); - agent.update_provider(provider, &session.id).await?; + agent + .update_provider(provider, ModelConfig::new("mock-model"), &session.id) + .await?; // Try to send a message - should trigger context limit, then recover via compaction let session_config = SessionConfig { diff --git a/crates/goose/tests/local_inference_integration.rs b/crates/goose/tests/local_inference_integration.rs index d52ab0c4bfc9..431aeb277438 100644 --- a/crates/goose/tests/local_inference_integration.rs +++ b/crates/goose/tests/local_inference_integration.rs @@ -28,8 +28,8 @@ fn test_model() -> String { #[tokio::test] #[ignore] async fn test_local_inference_stream_produces_output() { - let model_config = ModelConfig::new(test_model()).expect("valid model config"); - let provider = create("local", model_config.clone(), Vec::new()) + let model_config = ModelConfig::new(test_model()); + let provider = create("local", Vec::new()) .await .expect("provider creation should succeed"); @@ -70,10 +70,8 @@ async fn test_local_inference_stream_produces_output() { #[tokio::test] #[ignore] async fn test_local_inference_large_prompt() { - let model_config = ModelConfig::new(test_model()) - .expect("valid model config") - .with_max_tokens(Some(20)); - let provider = create("local", model_config.clone(), Vec::new()) + let model_config = ModelConfig::new(test_model()).with_max_tokens(Some(20)); + let provider = create("local", Vec::new()) .await .expect("provider creation should succeed"); @@ -137,8 +135,8 @@ async fn test_local_inference_vision_produces_output() { } }; - let model_config = ModelConfig::new(&model_id).expect("valid model config"); - let provider = create("local", model_config.clone(), Vec::new()) + let model_config = ModelConfig::new(&model_id); + let provider = create("local", Vec::new()) .await .expect("provider creation should succeed"); @@ -182,8 +180,8 @@ async fn test_local_inference_vision_produces_output() { #[tokio::test] #[ignore] async fn test_local_inference_vision_text_only_model_graceful() { - let model_config = ModelConfig::new(test_model()).expect("valid model config"); - let provider = create("local", model_config.clone(), Vec::new()) + let model_config = ModelConfig::new(test_model()); + let provider = create("local", Vec::new()) .await .expect("provider creation should succeed"); diff --git a/crates/goose/tests/local_inference_perf.rs b/crates/goose/tests/local_inference_perf.rs index 0fb4603dd5f9..bb9ba4fbe65f 100644 --- a/crates/goose/tests/local_inference_perf.rs +++ b/crates/goose/tests/local_inference_perf.rs @@ -24,10 +24,8 @@ fn test_model() -> String { #[tokio::test] #[ignore] async fn test_local_inference_cold_vs_warm() { - let model_config = ModelConfig::new(test_model()) - .expect("valid model config") - .with_max_tokens(Some(20)); - let provider = create("local", model_config.clone(), Vec::new()) + let model_config = ModelConfig::new(test_model()).with_max_tokens(Some(20)); + let provider = create("local", Vec::new()) .await .expect("provider creation should succeed"); diff --git a/crates/goose/tests/mcp_integration_test.rs b/crates/goose/tests/mcp_integration_test.rs index e7ee4c916e23..fc3ed1fd4ff2 100644 --- a/crates/goose/tests/mcp_integration_test.rs +++ b/crates/goose/tests/mcp_integration_test.rs @@ -40,14 +40,12 @@ struct Target { kind: Vec, } -#[derive(Clone)] -pub struct MockProvider { - pub model_config: ModelConfig, -} +#[derive(Clone, Default)] +pub struct MockProvider; impl MockProvider { - pub fn new(model_config: ModelConfig) -> Self { - Self { model_config } + pub fn new() -> Self { + Self } } @@ -61,11 +59,10 @@ impl ProviderDef for MockProvider { type Provider = Self; fn from_env( - model: ModelConfig, _extensions: Vec, _tls_config: Option, ) -> futures::future::BoxFuture<'static, anyhow::Result> { - Box::pin(async move { Ok(Self::new(model)) }) + Box::pin(async move { Ok(Self::new()) }) } } @@ -87,10 +84,6 @@ impl Provider for MockProvider { let usage = ProviderUsage::new("mock".to_string(), Usage::default()); Ok(stream_from_single_message(message, usage)) } - - fn get_model_config(&self) -> ModelConfig { - self.model_config.clone() - } } fn build_and_get_binary_path() -> PathBuf { @@ -184,7 +177,11 @@ async fn test_replayed_session( tool_calls: Vec, required_envs: Vec<&str>, ) { - std::env::set_var("GOOSE_MCP_CLIENT_VERSION", "0.0.0"); + let _env = env_lock::lock_env([ + ("GOOSE_MCP_CLIENT_VERSION", Some("0.0.0")), + ("GOOSE_PROVIDER", Some("openai")), + ("GOOSE_MODEL", Some("gpt-4o")), + ]); // Setup test file for developer extension tests let test_file_path = "/tmp/goose_test/goose.txt"; @@ -254,9 +251,9 @@ async fn test_replayed_session( available_tools: vec![], }; - let provider = Arc::new(tokio::sync::Mutex::new(Some(Arc::new(MockProvider { - model_config: ModelConfig::new("test-model").unwrap(), - }) as Arc))); + let provider = Arc::new(tokio::sync::Mutex::new(Some( + Arc::new(MockProvider::new()) as Arc + ))); let temp_dir = tempfile::tempdir().unwrap(); let session_manager = Arc::new(goose::session::SessionManager::new( temp_dir.path().to_path_buf(), diff --git a/crates/goose/tests/providers.rs b/crates/goose/tests/providers.rs index e9125822f50f..0231e60b1d15 100644 --- a/crates/goose/tests/providers.rs +++ b/crates/goose/tests/providers.rs @@ -101,6 +101,7 @@ struct ProviderFixture { expect_context_length_exceeded: bool, context_length_exceeded: usize, provider: Arc, + model_config: goose_providers::model::ModelConfig, agent: Agent, session_id: String, _mcp: McpFixture, @@ -234,11 +235,14 @@ impl ProviderFixture { let provider = create_with_named_model( &config.name.to_lowercase(), - config.model_name, vec![mcp_extension.clone(), developer_extension.clone()], ) .await .map_err(|e| anyhow::anyhow!("{}", e))?; + let model_config = goose::model_config::model_config_from_user_config( + &config.name.to_lowercase(), + config.model_name, + )?; let temp_dir = tempfile::tempdir()?; let session_manager = Arc::new(SessionManager::new(temp_dir.path().to_path_buf())); @@ -262,7 +266,9 @@ impl ProviderFixture { .await?; let session_id = session.id; expected_session_id.set(&session_id); - agent.update_provider(provider.clone(), &session_id).await?; + agent + .update_provider(provider.clone(), model_config.clone(), &session_id) + .await?; agent .add_extension(mcp_extension, &session_id) .await @@ -279,6 +285,7 @@ impl ProviderFixture { expect_context_length_exceeded: config.expect_context_length_exceeded, context_length_exceeded: config.context_length_exceeded, provider, + model_config, agent, session_id, _mcp: mcp, @@ -310,7 +317,7 @@ impl ProviderFixture { .build(); let message = Message::user().with_text(prompt); - let model_config = model_config.unwrap_or_else(|| self.provider.get_model_config()); + let model_config = model_config.unwrap_or_else(|| self.model_config.clone()); let (response1, _) = self .provider .complete( @@ -372,7 +379,7 @@ impl ProviderFixture { async fn test_basic_response(&self) -> Result<()> { let message = Message::user().with_text("Just say hello!"); - let model_config = self.provider.get_model_config(); + let model_config = self.model_config.clone(); let (response, _) = self .provider @@ -413,7 +420,7 @@ impl ProviderFixture { // "hello " ≈ 2 tokens across common tokenizers let large_message_content = "hello ".repeat(self.context_length_exceeded / 2); let messages = vec![Message::user().with_text(&large_message_content)]; - let model_config = self.provider.get_model_config(); + let model_config = self.model_config.clone(); let result = self .provider @@ -444,12 +451,9 @@ impl ProviderFixture { } async fn test_image_content_support(&self) -> Result<()> { - let image_config = match &self.image_model { - Some(model) => Some( - goose_providers::model::ModelConfig::new(model)?.with_canonical_limits(&self.name), - ), - None => None, - }; + let image_config = self.image_model.as_ref().map(|model| { + goose_providers::model::ModelConfig::new(model).with_canonical_limits(&self.name) + }); let response = self .tool_roundtrip( "Use the get_image tool and describe what you see in its result.", @@ -466,10 +470,10 @@ impl ProviderFixture { } async fn test_model_switch(&self) -> Result<()> { - let default = &self.provider.get_model_config().model_name; + let default = &self.model_config.model_name; let alt = self.model_switch_name.as_deref().unwrap(); let alt_config = - goose_providers::model::ModelConfig::new(alt)?.with_canonical_limits(&self.name); + goose_providers::model::ModelConfig::new(alt).with_canonical_limits(&self.name); let message = Message::user().with_text("Just say hello!"); let (response, _) = self @@ -505,7 +509,7 @@ impl ProviderFixture { println!("==================="); assert!(!models.is_empty()); - let resolved = &self.provider.get_model_config().model_name; + let resolved = &self.model_config.model_name; assert_ne!(resolved.as_str(), ACP_CURRENT_MODEL); assert!(models .iter() diff --git a/crates/goose/tests/session_id_propagation_test.rs b/crates/goose/tests/session_id_propagation_test.rs index 2d4f54e6d61a..dae559f1326a 100644 --- a/crates/goose/tests/session_id_propagation_test.rs +++ b/crates/goose/tests/session_id_propagation_test.rs @@ -42,8 +42,7 @@ fn create_test_provider(mock_server_url: &str) -> Box { None, ) .unwrap(); - let model = ModelConfig::new_or_fail("gpt-5-nano"); - Box::new(OpenAiProvider::new(api_client, model)) + Box::new(OpenAiProvider::new(api_client)) } async fn setup_mock_server() -> (MockServer, HeaderCapture, Box) { @@ -144,7 +143,7 @@ async fn setup_mock_server() -> (MockServer, HeaderCapture, Box) { async fn make_request(provider: &dyn Provider, session_id: &str) { let message = Message::user().with_text("test message"); - let model_config = provider.get_model_config(); + let model_config = ModelConfig::new("gpt-5-nano"); let _ = provider .complete( &model_config, diff --git a/crates/goose/tests/tetrate_streaming.rs b/crates/goose/tests/tetrate_streaming.rs index 111586a41c15..a065a10083ef 100644 --- a/crates/goose/tests/tetrate_streaming.rs +++ b/crates/goose/tests/tetrate_streaming.rs @@ -14,10 +14,7 @@ mod tetrate_streaming_tests { use super::*; async fn create_test_provider() -> Result { - // Create a test provider with the default model - let model_config = - ModelConfig::new("claude-3-5-sonnet-latest")?.with_canonical_limits("tetrate"); - TetrateProvider::from_env(model_config, None).await + TetrateProvider::from_env(None).await } #[tokio::test] @@ -27,7 +24,8 @@ mod tetrate_streaming_tests { let provider = create_test_provider().await?; let messages = vec![Message::user().with_text("Count from 1 to 5, one number at a time.")]; - let model_config = provider.get_model_config(); + let model_config = + ModelConfig::new("claude-3-5-sonnet-latest").with_canonical_limits("tetrate"); let mut stream = provider .stream( @@ -101,7 +99,8 @@ mod tetrate_streaming_tests { ); let messages = vec![Message::user().with_text("What's the weather in San Francisco?")]; - let model_config = provider.get_model_config(); + let model_config = + ModelConfig::new("claude-3-5-sonnet-latest").with_canonical_limits("tetrate"); let mut stream = provider .stream( @@ -151,7 +150,8 @@ mod tetrate_streaming_tests { // This might result in a very short or empty response let messages = vec![Message::user().with_text("")]; - let model_config = provider.get_model_config(); + let model_config = + ModelConfig::new("claude-3-5-sonnet-latest").with_canonical_limits("tetrate"); let mut stream = provider .stream( @@ -188,7 +188,8 @@ mod tetrate_streaming_tests { let messages = vec![Message::user().with_text( "Write a detailed 3-paragraph essay about the importance of streaming in modern APIs.", )]; - let model_config = provider.get_model_config(); + let model_config = + ModelConfig::new("claude-3-5-sonnet-latest").with_canonical_limits("tetrate"); let mut stream = provider .stream( @@ -246,12 +247,11 @@ mod tetrate_streaming_tests { // Test with invalid API key to ensure error handling works std::env::set_var("TETRATE_API_KEY", "invalid-key-for-testing"); - let model_config = - ModelConfig::new("claude-3-5-sonnet-latest")?.with_canonical_limits("tetrate"); - let provider = TetrateProvider::from_env(model_config, None).await?; + let provider = TetrateProvider::from_env(None).await?; let messages = vec![Message::user().with_text("Hello")]; - let model_config = provider.get_model_config(); + let model_config = + ModelConfig::new("claude-3-5-sonnet-latest").with_canonical_limits("tetrate"); let result = provider .stream( @@ -281,7 +281,8 @@ mod tetrate_streaming_tests { // Create multiple concurrent streams let messages1 = vec![Message::user().with_text("Say 'Stream 1'")]; let messages2 = vec![Message::user().with_text("Say 'Stream 2'")]; - let model_config = provider.get_model_config(); + let model_config = + ModelConfig::new("claude-3-5-sonnet-latest").with_canonical_limits("tetrate"); let stream1 = provider .stream( diff --git a/ui/desktop/openapi.json b/ui/desktop/openapi.json index 63d0a6ad5b84..1e7a991e423d 100644 --- a/ui/desktop/openapi.json +++ b/ui/desktop/openapi.json @@ -6910,6 +6910,11 @@ "type": "string", "description": "Display name for the provider in UIs" }, + "fast_model": { + "type": "string", + "description": "The name of a fast/cheap model to use for lightweight tasks (e.g. session naming,\ncompaction). When set, fast-path callers prefer this model over the main model.", + "nullable": true + }, "known_models": { "type": "array", "items": { diff --git a/ui/desktop/src/api/types.gen.ts b/ui/desktop/src/api/types.gen.ts index 26cb604a756f..b21738aed45b 100644 --- a/ui/desktop/src/api/types.gen.ts +++ b/ui/desktop/src/api/types.gen.ts @@ -1024,6 +1024,11 @@ export type ProviderMetadata = { * Display name for the provider in UIs */ display_name: string; + /** + * The name of a fast/cheap model to use for lightweight tasks (e.g. session naming, + * compaction). When set, fast-path callers prefer this model over the main model. + */ + fast_model?: string | null; /** * A list of currently known models with their capabilities */ From 20dc010895b1e1426c658b21f1a455733adbfd67 Mon Sep 17 00:00:00 2001 From: Douwe Osinga Date: Tue, 23 Jun 2026 21:31:25 -0400 Subject: [PATCH 09/12] Route MCP elicitations through tool streams (#9943) Signed-off-by: Douwe M Osinga Co-authored-by: Douwe M Osinga --- crates/goose/src/action_required_manager.rs | 255 +++++++++++--- crates/goose/src/agents/agent.rs | 62 ++-- crates/goose/src/agents/extension_manager.rs | 86 ++++- crates/goose/src/agents/mcp_client.rs | 317 ++++++++++++++++-- .../platform_extensions/code_execution.rs | 16 +- crates/goose/src/agents/tool_execution.rs | 8 +- crates/goose/src/session_context.rs | 1 + .../tests/mcp_replays/github-mcp-serverstdio | 2 +- ...ontextprotocol_server-everything@2026.1.14 | 10 +- ...14.4fastmcpruntests_fastmcp_test_server.py | 2 +- .../tests/mcp_replays/uvxmcp-server-fetch | 2 +- 11 files changed, 632 insertions(+), 129 deletions(-) diff --git a/crates/goose/src/action_required_manager.rs b/crates/goose/src/action_required_manager.rs index b76e0eaf8612..ce9a5c345aaa 100644 --- a/crates/goose/src/action_required_manager.rs +++ b/crates/goose/src/action_required_manager.rs @@ -1,9 +1,9 @@ use anyhow::Result; use serde_json::Value; -use std::collections::{HashMap, VecDeque}; +use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; -use tokio::sync::{Mutex, OwnedMutexGuard, RwLock}; +use tokio::sync::{mpsc, Mutex, OwnedMutexGuard, RwLock}; use tokio::time::timeout; use tracing::warn; use uuid::Uuid; @@ -46,14 +46,14 @@ impl PendingResponseClaim { pub(crate) struct ActionRequiredManager { pending: Arc>>>>, - queued_requests: Mutex>>, + action_required_senders: Mutex>>, } impl ActionRequiredManager { fn new() -> Self { Self { pending: Arc::new(RwLock::new(HashMap::new())), - queued_requests: Mutex::new(HashMap::new()), + action_required_senders: Mutex::new(HashMap::new()), } } @@ -66,6 +66,7 @@ impl ActionRequiredManager { pub(crate) async fn request_and_wait( &self, session_id: String, + tool_call_request_id: String, message: String, schema: Value, timeout_duration: Duration, @@ -87,12 +88,28 @@ impl ActionRequiredManager { MessageContent::action_required_elicitation(id.clone(), message, schema), ); - self.queued_requests + let sender = self + .action_required_senders .lock() .await - .entry(session_id) - .or_default() - .push_back(action_required_message); + .get(&(session_id.clone(), tool_call_request_id.clone())) + .cloned(); + + let Some(sender) = sender else { + self.pending.write().await.remove(&id); + return Err(anyhow::anyhow!( + "Tool call request not found for elicitation: {}", + tool_call_request_id + )); + }; + + if sender.send(action_required_message).await.is_err() { + self.pending.write().await.remove(&id); + return Err(anyhow::anyhow!( + "Tool call action-required stream closed: {}", + tool_call_request_id + )); + } let result = self .wait_for_response(&id, pending_request, rx, timeout_duration) @@ -179,13 +196,39 @@ impl ActionRequiredManager { } } - pub(crate) async fn drain_requests_for_session(&self, session_id: &str) -> Vec { - self.queued_requests + pub(crate) async fn register_action_required_stream( + &self, + session_id: String, + tool_call_request_id: String, + ) -> mpsc::Receiver { + let (tx, rx) = mpsc::channel(8); + self.action_required_senders + .lock() + .await + .insert((session_id, tool_call_request_id), tx); + rx + } + + pub(crate) async fn has_action_required_stream( + &self, + session_id: &str, + tool_call_request_id: &str, + ) -> bool { + self.action_required_senders .lock() .await - .remove(session_id) - .map(|queue| queue.into_iter().collect()) - .unwrap_or_default() + .contains_key(&(session_id.to_string(), tool_call_request_id.to_string())) + } + + pub(crate) async fn unregister_action_required_stream( + &self, + session_id: &str, + tool_call_request_id: &str, + ) { + self.action_required_senders + .lock() + .await + .remove(&(session_id.to_string(), tool_call_request_id.to_string())); } } @@ -205,32 +248,26 @@ mod tests { } } - async fn wait_for_elicitation_messages( - manager: &ActionRequiredManager, - session_id: &str, - ) -> Vec { - tokio::time::timeout(Duration::from_secs(1), async { - loop { - let messages = manager.drain_requests_for_session(session_id).await; - if !messages.is_empty() { - return messages; - } - tokio::task::yield_now().await; - } - }) - .await - .unwrap_or_else(|_| panic!("timed out waiting for elicitation message for {session_id}")) + async fn recv_elicitation_message(rx: &mut mpsc::Receiver) -> Message { + tokio::time::timeout(Duration::from_secs(1), rx.recv()) + .await + .expect("timed out waiting for elicitation message") + .expect("action-required stream closed") } #[tokio::test] async fn wrong_session_does_not_consume_pending_response() { let manager = Arc::new(ActionRequiredManager::new()); + let mut action_required_rx = manager + .register_action_required_stream("session-a".to_string(), "tool-call-a".to_string()) + .await; let waiter = { let manager = manager.clone(); tokio::spawn(async move { manager .request_and_wait( "session-a".to_string(), + "tool-call-a".to_string(), "Need input".to_string(), json!({ "type": "object" }), Duration::from_secs(5), @@ -239,9 +276,8 @@ mod tests { }) }; - let messages = wait_for_elicitation_messages(&manager, "session-a").await; - assert_eq!(messages.len(), 1); - let request_id = elicitation_id(&messages[0]); + let message = recv_elicitation_message(&mut action_required_rx).await; + let request_id = elicitation_id(&message); let err = match manager.claim_response("session-b", &request_id).await { Ok(_) => panic!("wrong session should not claim pending response"), @@ -264,14 +300,21 @@ mod tests { } #[tokio::test] - async fn drains_only_requested_session() { + async fn streams_only_requested_tool_call() { let manager = Arc::new(ActionRequiredManager::new()); + let mut stream_a = manager + .register_action_required_stream("session-a".to_string(), "tool-call-a".to_string()) + .await; + let mut stream_b = manager + .register_action_required_stream("session-b".to_string(), "tool-call-b".to_string()) + .await; let waiter_a = { let manager = manager.clone(); tokio::spawn(async move { manager .request_and_wait( "session-a".to_string(), + "tool-call-a".to_string(), "Need input A".to_string(), json!({ "type": "object" }), Duration::from_secs(5), @@ -285,6 +328,7 @@ mod tests { manager .request_and_wait( "session-b".to_string(), + "tool-call-b".to_string(), "Need input B".to_string(), json!({ "type": "object" }), Duration::from_secs(5), @@ -293,16 +337,78 @@ mod tests { }) }; - let session_a_messages = wait_for_elicitation_messages(&manager, "session-a").await; - assert_eq!(session_a_messages.len(), 1); - let request_id_a = elicitation_id(&session_a_messages[0]); + let message_a = recv_elicitation_message(&mut stream_a).await; + let request_id_a = elicitation_id(&message_a); + assert!(stream_a.try_recv().is_err()); - let empty_messages = manager.drain_requests_for_session("session-a").await; - assert!(empty_messages.is_empty()); + let message_b = recv_elicitation_message(&mut stream_b).await; + let request_id_b = elicitation_id(&message_b); - let session_b_messages = wait_for_elicitation_messages(&manager, "session-b").await; - assert_eq!(session_b_messages.len(), 1); - let request_id_b = elicitation_id(&session_b_messages[0]); + manager + .claim_response("session-a", &request_id_a) + .await + .unwrap() + .submit(ElicitationOutcome::Accept(json!({ "answer": "a" }))) + .unwrap(); + manager + .claim_response("session-b", &request_id_b) + .await + .unwrap() + .submit(ElicitationOutcome::Accept(json!({ "answer": "b" }))) + .unwrap(); + + assert_eq!( + waiter_a.await.unwrap().unwrap(), + ElicitationOutcome::Accept(json!({ "answer": "a" })) + ); + assert_eq!( + waiter_b.await.unwrap().unwrap(), + ElicitationOutcome::Accept(json!({ "answer": "b" })) + ); + } + + #[tokio::test] + async fn streams_are_namespaced_by_session() { + let manager = Arc::new(ActionRequiredManager::new()); + let mut stream_a = manager + .register_action_required_stream("session-a".to_string(), "tool-call-a".to_string()) + .await; + let mut stream_b = manager + .register_action_required_stream("session-b".to_string(), "tool-call-a".to_string()) + .await; + let waiter_a = { + let manager = manager.clone(); + tokio::spawn(async move { + manager + .request_and_wait( + "session-a".to_string(), + "tool-call-a".to_string(), + "Need input A".to_string(), + json!({ "type": "object" }), + Duration::from_secs(5), + ) + .await + }) + }; + let waiter_b = { + let manager = manager.clone(); + tokio::spawn(async move { + manager + .request_and_wait( + "session-b".to_string(), + "tool-call-a".to_string(), + "Need input B".to_string(), + json!({ "type": "object" }), + Duration::from_secs(5), + ) + .await + }) + }; + + let message_a = recv_elicitation_message(&mut stream_a).await; + let request_id_a = elicitation_id(&message_a); + let message_b = recv_elicitation_message(&mut stream_b).await; + let request_id_b = elicitation_id(&message_b); manager .claim_response("session-a", &request_id_a) @@ -330,12 +436,16 @@ mod tests { #[tokio::test] async fn claimed_response_can_complete_after_timeout_deadline() { let manager = Arc::new(ActionRequiredManager::new()); + let mut action_required_rx = manager + .register_action_required_stream("session-a".to_string(), "tool-call-a".to_string()) + .await; let waiter = { let manager = manager.clone(); tokio::spawn(async move { manager .request_and_wait( "session-a".to_string(), + "tool-call-a".to_string(), "Need input".to_string(), json!({ "type": "object" }), Duration::from_millis(25), @@ -344,9 +454,8 @@ mod tests { }) }; - let messages = wait_for_elicitation_messages(&manager, "session-a").await; - assert_eq!(messages.len(), 1); - let request_id = elicitation_id(&messages[0]); + let message = recv_elicitation_message(&mut action_required_rx).await; + let request_id = elicitation_id(&message); let claim = manager .claim_response("session-a", &request_id) @@ -368,12 +477,19 @@ mod tests { #[tokio::test] async fn request_and_wait_returns_decline_and_cancel_actions() { let manager = Arc::new(ActionRequiredManager::new()); + let mut decline_rx = manager + .register_action_required_stream("session-a".to_string(), "tool-call-a".to_string()) + .await; + let mut cancel_rx = manager + .register_action_required_stream("session-b".to_string(), "tool-call-b".to_string()) + .await; let decline_waiter = { let manager = manager.clone(); tokio::spawn(async move { manager .request_and_wait( "session-a".to_string(), + "tool-call-a".to_string(), "Need input A".to_string(), json!({ "type": "object" }), Duration::from_secs(5), @@ -387,6 +503,7 @@ mod tests { manager .request_and_wait( "session-b".to_string(), + "tool-call-b".to_string(), "Need input B".to_string(), json!({ "type": "object" }), Duration::from_secs(5), @@ -395,10 +512,10 @@ mod tests { }) }; - let decline_messages = wait_for_elicitation_messages(&manager, "session-a").await; - let decline_request_id = elicitation_id(&decline_messages[0]); - let cancel_messages = wait_for_elicitation_messages(&manager, "session-b").await; - let cancel_request_id = elicitation_id(&cancel_messages[0]); + let decline_message = recv_elicitation_message(&mut decline_rx).await; + let decline_request_id = elicitation_id(&decline_message); + let cancel_message = recv_elicitation_message(&mut cancel_rx).await; + let cancel_request_id = elicitation_id(&cancel_message); manager .claim_response("session-a", &decline_request_id) @@ -422,4 +539,48 @@ mod tests { ElicitationOutcome::Cancel ); } + + #[tokio::test] + async fn missing_tool_call_stream_errors() { + let manager = Arc::new(ActionRequiredManager::new()); + + let result = manager + .request_and_wait( + "session-a".to_string(), + "missing-tool-call".to_string(), + "Need input".to_string(), + json!({ "type": "object" }), + Duration::from_secs(5), + ) + .await; + + let err = result.expect_err("request should fail without a registered stream"); + assert!(err + .to_string() + .contains("Tool call request not found for elicitation")); + } + + #[tokio::test] + async fn closed_tool_call_stream_errors() { + let manager = Arc::new(ActionRequiredManager::new()); + let rx = manager + .register_action_required_stream("session-a".to_string(), "tool-call-a".to_string()) + .await; + drop(rx); + + let result = manager + .request_and_wait( + "session-a".to_string(), + "tool-call-a".to_string(), + "Need input".to_string(), + json!({ "type": "object" }), + Duration::from_secs(5), + ) + .await; + + let err = result.expect_err("request should fail when stream is closed"); + assert!(err + .to_string() + .contains("Tool call action-required stream closed")); + } } diff --git a/crates/goose/src/agents/agent.rs b/crates/goose/src/agents/agent.rs index 4665d76b3f02..3e62fd6ae62e 100644 --- a/crates/goose/src/agents/agent.rs +++ b/crates/goose/src/agents/agent.rs @@ -16,7 +16,7 @@ use super::mcp_client::GooseMcpHostInfo; use super::platform_tools; use super::tool_confirmation_router::ToolConfirmationRouter; use super::tool_execution::{ToolCallResult, CHAT_MODE_TOOL_SKIPPED_RESPONSE, DECLINED_RESPONSE}; -use crate::action_required_manager::{ActionRequiredManager, ElicitationOutcome}; +use crate::action_required_manager::ElicitationOutcome; use crate::agents::extension::{ExtensionConfig, ExtensionResult, ToolInfo}; use crate::agents::extension_manager::{ get_parameter_names, ExtensionManager, ExtensionManagerCapabilities, @@ -274,6 +274,7 @@ impl Default for Agent { } pub enum ToolStreamItem { + ActionRequired(Message), Message(ServerNotification), Result(T), } @@ -285,17 +286,22 @@ pub type ToolStream = // final result of the tool call. MCP notifications are not request-scoped, but // this lets us capture all notifications emitted during the tool call for // simpler consumption -pub fn tool_stream(rx: S, done: F) -> ToolStream +pub fn tool_stream(rx: S, action_required_rx: A, done: F) -> ToolStream where S: Stream + Send + Unpin + 'static, + A: Stream + Send + Unpin + 'static, F: Future> + Send + 'static, { Box::pin(async_stream::stream! { tokio::pin!(done); let mut rx = rx; + let mut action_required_rx = action_required_rx; loop { tokio::select! { + Some(msg) = action_required_rx.next() => { + yield ToolStreamItem::ActionRequired(msg); + } Some(msg) = rx.next() => { yield ToolStreamItem::Message(msg); } @@ -572,6 +578,7 @@ impl Agent { ToolCallResult { notification_stream: result.notification_stream, + action_required_stream: result.action_required_stream, result: Box::new(fut.boxed()), } } @@ -645,24 +652,6 @@ impl Agent { | RetryResult::SuccessChecksPassed => Ok(false), } } - async fn drain_elicitation_messages(&self, session_id: &str) -> Vec { - let mut messages = Vec::new(); - let manager = self.config.session_manager.clone(); - for mut elicitation_message in ActionRequiredManager::global() - .drain_requests_for_session(session_id) - .await - { - if elicitation_message.id.is_none() { - elicitation_message = elicitation_message.with_generated_id(); - } - if let Err(e) = manager.add_message(session_id, &elicitation_message).await { - warn!("Failed to save elicitation message to session: {}", e); - } - messages.push(elicitation_message); - } - messages - } - async fn load_project_instructions(&self, session: &Session) -> Option { let project_id = session.project_id.as_deref()?; let entry = crate::sources::read_project(project_id).ok()?; @@ -783,11 +772,16 @@ impl Agent { result .notification_stream .unwrap_or_else(|| Box::new(stream::empty())), + result + .action_required_stream + .unwrap_or_else(|| Box::new(stream::empty())), result.result, ), - Err(e) => { - tool_stream(Box::new(stream::empty()), futures::future::ready(Err(e))) - } + Err(e) => tool_stream( + Box::new(stream::empty()), + Box::new(stream::empty()), + futures::future::ready(Err(e)), + ), }, )); } @@ -2152,10 +2146,6 @@ impl Agent { break; } - for msg in self.drain_elicitation_messages(&session_config.id).await { - yield AgentEvent::Message(msg); - } - tokio::select! { biased; @@ -2163,6 +2153,15 @@ impl Agent { match tool_item { Some((request_id, item)) => { match item { + ToolStreamItem::ActionRequired(mut msg) => { + if msg.id.is_none() { + msg = msg.with_generated_id(); + } + if let Err(e) = session_manager.add_message(&session_config.id, &msg).await { + warn!("Failed to save elicitation message to session: {}", e); + } + yield AgentEvent::Message(msg); + } ToolStreamItem::Result(output) => { if let Ok(ref call_result) = output { if let Some(ref meta) = call_result.meta { @@ -2200,17 +2199,10 @@ impl Agent { } } - _ = tokio::time::sleep(std::time::Duration::from_millis(100)) => { - // Continue loop to drain elicitation messages - } + _ = tokio::time::sleep(std::time::Duration::from_millis(100)) => {} } } - // check for remaining elicitation messages after all tools complete - for msg in self.drain_elicitation_messages(&session_config.id).await { - yield AgentEvent::Message(msg); - } - if all_install_successful && !enable_extension_request_ids.is_empty() { if let Err(e) = self.save_extension_state(&session_config).await { warn!("Failed to save extension state after runtime changes: {}", e); diff --git a/crates/goose/src/agents/extension_manager.rs b/crates/goose/src/agents/extension_manager.rs index ec2fcc489534..fc52e7d6c182 100644 --- a/crates/goose/src/agents/extension_manager.rs +++ b/crates/goose/src/agents/extension_manager.rs @@ -2,6 +2,7 @@ use anyhow::Result; use axum::http::{HeaderMap, HeaderName, HeaderValue}; use chrono::{DateTime, Utc}; use futures::stream::{FuturesUnordered, StreamExt}; +use futures::Stream; use futures::{future, FutureExt}; use once_cell::sync::Lazy; use rmcp::service::{ClientInitializeError, ServiceError}; @@ -13,9 +14,11 @@ use rmcp::transport::{ }; use std::collections::HashMap; use std::path::PathBuf; +use std::pin::Pin; use std::process::Stdio; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; +use std::task::{Context, Poll}; use std::time::Duration; use tempfile::{tempdir, TempDir}; use tokio::io::AsyncReadExt; @@ -32,6 +35,7 @@ use super::extension::{ }; use super::tool_execution::{ToolCallContext, ToolCallResult}; use super::types::SharedProvider; +use crate::action_required_manager::ActionRequiredManager; use crate::agents::extension::{Envs, ProcessExit}; use crate::agents::extension_malware_check; use crate::agents::mcp_client::{ @@ -54,6 +58,49 @@ use serde_json::Value; type McpClientBox = Arc; +struct ActionRequiredStream { + inner: ReceiverStream, + session_id: String, + tool_call_request_id: String, +} + +impl ActionRequiredStream { + fn new( + receiver: tokio::sync::mpsc::Receiver, + session_id: String, + tool_call_request_id: String, + ) -> Self { + Self { + inner: ReceiverStream::new(receiver), + session_id, + tool_call_request_id, + } + } +} + +impl Stream for ActionRequiredStream { + type Item = crate::conversation::message::Message; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inner).poll_next(cx) + } +} + +impl Drop for ActionRequiredStream { + fn drop(&mut self) { + let session_id = self.session_id.clone(); + let tool_call_request_id = self.tool_call_request_id.clone(); + let Ok(handle) = tokio::runtime::Handle::try_current() else { + return; + }; + handle.spawn(async move { + ActionRequiredManager::global() + .unregister_action_required_stream(&session_id, &tool_call_request_id) + .await; + }); + } +} + static RE_ENV_BRACES: Lazy = Lazy::new(|| regex::Regex::new(r"\$\{\s*([A-Za-z_][A-Za-z0-9_]*)\s*\}").expect("valid regex")); @@ -1733,11 +1780,33 @@ impl ExtensionManager { let client = resolved.client.clone(); let hydration_client = client.clone(); let notifications_receiver = client.subscribe().await; + let session_id = ctx.session_id.clone(); + let action_required_tool_call_request_id = ctx.tool_call_request_id.clone(); + let action_required_receiver = + if let Some(tool_call_request_id) = action_required_tool_call_request_id.clone() { + if ActionRequiredManager::global() + .has_action_required_stream(&session_id, &tool_call_request_id) + .await + { + None + } else { + let registered_tool_call_request_id = tool_call_request_id.clone(); + let receiver = ActionRequiredManager::global() + .register_action_required_stream(session_id.clone(), tool_call_request_id) + .await; + Some(( + receiver, + session_id.clone(), + registered_tool_call_request_id, + )) + } + } else { + None + }; let actual_tool_name = resolved.actual_tool_name.clone(); let resolved_tool = resolved; let should_hydrate_mcp_app = self.host_supports_mcp_apps(); let read_cancellation_token = cancellation_token.clone(); - let session_id = ctx.session_id.clone(); let owned_ctx = ToolCallContext::new( ctx.session_id.clone(), ctx.working_dir.clone(), @@ -1751,7 +1820,7 @@ impl ExtensionManager { owned_ctx.session_id, owned_ctx.working_dir, ); - let mut result = client + let call_result = client .call_tool(&owned_ctx, &actual_tool_name, arguments, cancellation_token) .await .map_err(|e| match e { @@ -1759,7 +1828,9 @@ impl ExtensionManager { _ => { ErrorData::new(ErrorCode::INTERNAL_ERROR, e.to_string(), e.maybe_to_value()) } - })?; + }); + + let mut result = call_result?; remove_untrusted_mcp_app_meta(&mut result); @@ -1782,6 +1853,15 @@ impl ExtensionManager { Ok(ToolCallResult { result: Box::new(fut.boxed()), notification_stream: Some(Box::new(ReceiverStream::new(notifications_receiver))), + action_required_stream: action_required_receiver.map( + |(rx, session_id, tool_call_request_id)| { + Box::new(ActionRequiredStream::new( + rx, + session_id, + tool_call_request_id, + )) as _ + }, + ), }) } diff --git a/crates/goose/src/agents/mcp_client.rs b/crates/goose/src/agents/mcp_client.rs index 913a7a816fe0..919ca3fa4e2f 100644 --- a/crates/goose/src/agents/mcp_client.rs +++ b/crates/goose/src/agents/mcp_client.rs @@ -1,7 +1,7 @@ use crate::action_required_manager::{ActionRequiredManager, ElicitationOutcome}; use crate::agents::tool_execution::ToolCallContext; use crate::agents::types::SharedProvider; -use crate::session_context::{SESSION_ID_HEADER, WORKING_DIR_HEADER}; +use crate::session_context::{SESSION_ID_HEADER, TOOL_CALL_REQUEST_ID_HEADER, WORKING_DIR_HEADER}; use rmcp::model::{ CreateElicitationRequestParams, CreateElicitationResult, ElicitationAction, ErrorCode, ExtensionCapabilities, Extensions, JsonObject, ListRootsResult, LoggingMessageNotification, @@ -26,7 +26,9 @@ use rmcp::{ ClientHandler, ErrorData, Peer, RoleClient, ServiceError, ServiceExt, }; use serde_json::Value; -use std::{path::PathBuf, sync::Arc, time::Duration}; +use std::{ + collections::HashMap, path::PathBuf, sync::Arc, sync::Mutex as StdMutex, time::Duration, +}; use tokio::sync::{ mpsc::{self, Sender}, Mutex, @@ -148,10 +150,34 @@ pub trait McpClientTrait: Send + Sync { } } +struct ActiveToolCallGuard { + active_tool_calls: Arc>>>, + session_id: String, + tool_call_request_id: String, +} + +impl Drop for ActiveToolCallGuard { + fn drop(&mut self) { + let mut active_tool_calls = self + .active_tool_calls + .lock() + .expect("active_tool_calls mutex poisoned"); + if let Some(calls) = active_tool_calls.get_mut(&self.session_id) { + if let Some(pos) = calls.iter().position(|id| id == &self.tool_call_request_id) { + calls.remove(pos); + } + if calls.is_empty() { + active_tool_calls.remove(&self.session_id); + } + } + } +} + pub struct GooseClient { notification_handlers: Arc>>>, provider: SharedProvider, session_id: Mutex>, + active_tool_calls: Arc>>>, client_name: String, capabilities: GooseMcpClientCapabilities, working_dir: Arc>, @@ -169,6 +195,7 @@ impl GooseClient { notification_handlers: handlers, provider, session_id: Mutex::new(None), + active_tool_calls: Arc::new(StdMutex::new(HashMap::new())), client_name, capabilities, working_dir: Arc::new(tokio::sync::RwLock::new(working_dir)), @@ -207,6 +234,62 @@ impl GooseClient { .map(|value| value.to_string()) } + fn tool_call_request_id_from_extensions(extensions: &Extensions) -> Option { + let meta = extensions.get::()?; + meta.0 + .iter() + .find(|(key, _)| key.eq_ignore_ascii_case(TOOL_CALL_REQUEST_ID_HEADER)) + .and_then(|(_, value)| value.as_str()) + .map(|value| value.to_string()) + } + + fn register_active_tool_call( + &self, + session_id: &str, + tool_call_request_id: &str, + ) -> ActiveToolCallGuard { + self.active_tool_calls + .lock() + .expect("active_tool_calls mutex poisoned") + .entry(session_id.to_string()) + .or_default() + .push(tool_call_request_id.to_string()); + ActiveToolCallGuard { + active_tool_calls: self.active_tool_calls.clone(), + session_id: session_id.to_string(), + tool_call_request_id: tool_call_request_id.to_string(), + } + } + + fn resolve_tool_call_request_id( + &self, + session_id: &str, + extensions: &Extensions, + ) -> Result { + if let Some(tool_call_request_id) = Self::tool_call_request_id_from_extensions(extensions) { + return Ok(tool_call_request_id); + } + + let active_tool_calls = self + .active_tool_calls + .lock() + .expect("active_tool_calls mutex poisoned"); + match active_tool_calls.get(session_id).map(Vec::as_slice) { + Some([tool_call_request_id]) => Ok(tool_call_request_id.clone()), + Some(calls) if calls.len() > 1 => Err(ErrorData::new( + ErrorCode::INTERNAL_ERROR, + "Cannot correlate elicitation request: multiple tool calls are active and the \ + server did not echo the tool call request id", + None, + )), + _ => Err(ErrorData::new( + ErrorCode::INTERNAL_ERROR, + "Could not resolve tool call request id for elicitation request", + None, + )), + } + } + fn resolved_extensions(&self) -> ExtensionCapabilities { if let Some(host_info) = &self.capabilities.host_info { if host_info.explicit_extensions { @@ -396,6 +479,8 @@ impl ClientHandler for GooseClient { None, ) })?; + let tool_call_request_id = + self.resolve_tool_call_request_id(&session_id, &context.extensions)?; let (message, schema_value) = match &request { CreateElicitationRequestParams::FormElicitationParams { @@ -418,7 +503,13 @@ impl ClientHandler for GooseClient { }; ActionRequiredManager::global() - .request_and_wait(session_id, message, schema_value, Duration::from_secs(300)) + .request_and_wait( + session_id, + tool_call_request_id, + message, + schema_value, + Duration::from_secs(300), + ) .await .map(|response| match response { ElicitationOutcome::Accept(user_data) => { @@ -548,19 +639,34 @@ impl McpClient { &self, session_id: &str, working_dir: Option<&str>, + tool_call_request_id: Option<&str>, request: ClientRequest, cancel_token: CancellationToken, ) -> Result { - let request = inject_session_context_into_request(request, Some(session_id), working_dir); + let request = inject_session_context_into_request( + request, + Some(session_id), + working_dir, + tool_call_request_id, + ); + let active_tool_call = tool_call_request_id.filter(|id| !id.is_empty()); // The inner mutex is held only for the send; the actual response wait - // happens outside the lock so concurrent calls can overlap. - let handle = { + // happens outside the lock so concurrent calls can overlap. The guard + // unregisters the active tool call on drop, covering cancellation and + // dropped reply streams as well as normal completion. + let (handle, _active_tool_call_guard) = { let client = self.client.lock().await; client.service().set_session_id(session_id).await; - client + let guard = active_tool_call.map(|tool_call_request_id| { + client + .service() + .register_active_tool_call(session_id, tool_call_request_id) + }); + let handle = client .send_cancellable_request(request, PeerRequestOptions::no_options()) - .await - }?; + .await?; + (handle, guard) + }; await_response(handle, self.timeout, &cancel_token).await } @@ -616,6 +722,7 @@ impl McpClientTrait for McpClient { .send_request_with_context( session_id, None, + None, ClientRequest::ListResourcesRequest(RequestOptionalParam::with_param( PaginatedRequestParams::default().with_cursor(cursor), )), @@ -639,6 +746,7 @@ impl McpClientTrait for McpClient { .send_request_with_context( session_id, None, + None, ClientRequest::ReadResourceRequest(Request::new(ReadResourceRequestParams::new( uri.to_string(), ))), @@ -662,6 +770,7 @@ impl McpClientTrait for McpClient { .send_request_with_context( session_id, None, + None, ClientRequest::ListToolsRequest(RequestOptionalParam::with_param( PaginatedRequestParams::default().with_cursor(cursor), )), @@ -692,6 +801,7 @@ impl McpClientTrait for McpClient { .send_request_with_context( &ctx.session_id, ctx.working_dir_str(), + ctx.tool_call_request_id.as_deref(), request, cancel_token, ) @@ -713,6 +823,7 @@ impl McpClientTrait for McpClient { .send_request_with_context( session_id, None, + None, ClientRequest::ListPromptsRequest(RequestOptionalParam::with_param( PaginatedRequestParams::default().with_cursor(cursor), )), @@ -745,6 +856,7 @@ impl McpClientTrait for McpClient { .send_request_with_context( session_id, None, + None, ClientRequest::GetPromptRequest(Request::new(params)), cancel_token, ) @@ -773,9 +885,11 @@ fn inject_session_context_into_extensions( mut extensions: Extensions, session_id: Option<&str>, working_dir: Option<&str>, + tool_call_request_id: Option<&str>, ) -> Extensions { let session_id = session_id.filter(|id| !id.is_empty()); let working_dir = working_dir.filter(|dir| !dir.is_empty()); + let tool_call_request_id = tool_call_request_id.filter(|id| !id.is_empty()); let mut meta_map = extensions .get::() .map(|meta| meta.0.clone()) @@ -783,7 +897,9 @@ fn inject_session_context_into_extensions( // JsonObject is case-sensitive, so we use retain for case-insensitive removal meta_map.retain(|k, _| { - !k.eq_ignore_ascii_case(SESSION_ID_HEADER) && !k.eq_ignore_ascii_case(WORKING_DIR_HEADER) + !k.eq_ignore_ascii_case(SESSION_ID_HEADER) + && !k.eq_ignore_ascii_case(WORKING_DIR_HEADER) + && !k.eq_ignore_ascii_case(TOOL_CALL_REQUEST_ID_HEADER) }); if let Some(session_id) = session_id { @@ -800,6 +916,13 @@ fn inject_session_context_into_extensions( ); } + if let Some(tool_call_request_id) = tool_call_request_id { + meta_map.insert( + TOOL_CALL_REQUEST_ID_HEADER.to_string(), + Value::String(tool_call_request_id.to_string()), + ); + } + extensions.insert(Meta(meta_map)); extensions } @@ -808,36 +931,61 @@ fn inject_session_context_into_request( request: ClientRequest, session_id: Option<&str>, working_dir: Option<&str>, + tool_call_request_id: Option<&str>, ) -> ClientRequest { match request { ClientRequest::ListResourcesRequest(mut req) => { - req.extensions = - inject_session_context_into_extensions(req.extensions, session_id, working_dir); + req.extensions = inject_session_context_into_extensions( + req.extensions, + session_id, + working_dir, + None, + ); ClientRequest::ListResourcesRequest(req) } ClientRequest::ReadResourceRequest(mut req) => { - req.extensions = - inject_session_context_into_extensions(req.extensions, session_id, working_dir); + req.extensions = inject_session_context_into_extensions( + req.extensions, + session_id, + working_dir, + None, + ); ClientRequest::ReadResourceRequest(req) } ClientRequest::ListToolsRequest(mut req) => { - req.extensions = - inject_session_context_into_extensions(req.extensions, session_id, working_dir); + req.extensions = inject_session_context_into_extensions( + req.extensions, + session_id, + working_dir, + None, + ); ClientRequest::ListToolsRequest(req) } ClientRequest::CallToolRequest(mut req) => { - req.extensions = - inject_session_context_into_extensions(req.extensions, session_id, working_dir); + req.extensions = inject_session_context_into_extensions( + req.extensions, + session_id, + working_dir, + tool_call_request_id, + ); ClientRequest::CallToolRequest(req) } ClientRequest::ListPromptsRequest(mut req) => { - req.extensions = - inject_session_context_into_extensions(req.extensions, session_id, working_dir); + req.extensions = inject_session_context_into_extensions( + req.extensions, + session_id, + working_dir, + None, + ); ClientRequest::ListPromptsRequest(req) } ClientRequest::GetPromptRequest(mut req) => { - req.extensions = - inject_session_context_into_extensions(req.extensions, session_id, working_dir); + req.extensions = inject_session_context_into_extensions( + req.extensions, + session_id, + working_dir, + None, + ); ClientRequest::GetPromptRequest(req) } other => other, @@ -953,7 +1101,7 @@ mod tests { } let extensions = - inject_session_context_into_extensions(Extensions::new(), ext_session, None); + inject_session_context_into_extensions(Extensions::new(), ext_session, None, None); let resolved = client.resolve_session_id(&extensions).await; @@ -962,6 +1110,83 @@ mod tests { }); } + #[test] + fn test_resolve_tool_call_request_id_from_extensions() { + let client = new_client(GoosePlatform::GooseCli); + let _guard = client.register_active_tool_call("session-a", "active-tool-call"); + let extensions = inject_session_context_into_extensions( + Extensions::new(), + Some("session-a"), + None, + Some("extension-tool-call"), + ); + + let resolved = client + .resolve_tool_call_request_id("session-a", &extensions) + .unwrap(); + + assert_eq!(resolved, "extension-tool-call"); + } + + #[test] + fn test_resolve_tool_call_request_id_from_active_call() { + let client = new_client(GoosePlatform::GooseCli); + let _guard = client.register_active_tool_call("session-a", "active-tool-call"); + + let resolved = client + .resolve_tool_call_request_id("session-a", &Extensions::new()) + .unwrap(); + + assert_eq!(resolved, "active-tool-call"); + } + + #[test] + fn test_resolve_tool_call_request_id_errors_when_calls_overlap() { + let client = new_client(GoosePlatform::GooseCli); + let _guard_a = client.register_active_tool_call("session-a", "active-tool-call-a"); + let _guard_b = client.register_active_tool_call("session-a", "active-tool-call-b"); + + let error = client + .resolve_tool_call_request_id("session-a", &Extensions::new()) + .expect_err("ambiguous elicitation should not resolve to an arbitrary call"); + + assert_eq!(error.code, ErrorCode::INTERNAL_ERROR); + } + + #[test] + fn test_resolve_tool_call_request_id_prefers_echoed_id_while_calls_overlap() { + let client = new_client(GoosePlatform::GooseCli); + let _guard_a = client.register_active_tool_call("session-a", "active-tool-call-a"); + let _guard_b = client.register_active_tool_call("session-a", "active-tool-call-b"); + let extensions = inject_session_context_into_extensions( + Extensions::new(), + Some("session-a"), + None, + Some("active-tool-call-a"), + ); + + let resolved = client + .resolve_tool_call_request_id("session-a", &extensions) + .unwrap(); + + assert_eq!(resolved, "active-tool-call-a"); + } + + #[test] + fn test_dropping_guard_unregisters_active_tool_call() { + let client = new_client(GoosePlatform::GooseCli); + let guard_a = client.register_active_tool_call("session-a", "active-tool-call-a"); + let _guard_b = client.register_active_tool_call("session-a", "active-tool-call-b"); + + drop(guard_a); + + let resolved = client + .resolve_tool_call_request_id("session-a", &Extensions::new()) + .unwrap(); + + assert_eq!(resolved, "active-tool-call-b"); + } + #[test_case(list_resources_request; "list_resources")] #[test_case(read_resource_request; "read_resource")] #[test_case(list_tools_request; "list_tools")] @@ -980,7 +1205,7 @@ mod tests { ); let request = request_builder(extensions); - let request = inject_session_context_into_request(request, Some(session_id), None); + let request = inject_session_context_into_request(request, Some(session_id), None, None); let extensions = request_extensions(&request).expect("request should have extensions"); let meta = extensions .get::() @@ -994,13 +1219,20 @@ mod tests { meta.0.get("other-key"), Some(&Value::String("preserve-me".to_string())) ); + if matches!(request, ClientRequest::CallToolRequest(_)) { + assert!(!meta.0.contains_key(TOOL_CALL_REQUEST_ID_HEADER)); + } } #[test] fn test_session_id_in_mcp_meta() { let session_id = "test-session-789"; - let extensions = - inject_session_context_into_extensions(Default::default(), Some(session_id), None); + let extensions = inject_session_context_into_extensions( + Default::default(), + Some(session_id), + None, + None, + ); let mcp_meta = extensions.get::().unwrap(); assert_eq!( @@ -1052,12 +1284,43 @@ mod tests { .unwrap(), ); - let extensions = inject_session_context_into_extensions(extensions, session_id, None); + let extensions = inject_session_context_into_extensions(extensions, session_id, None, None); let mcp_meta = extensions.get::().unwrap(); assert_eq!(&mcp_meta.0, expected_meta.as_object().unwrap()); } + #[test] + fn test_tool_call_request_id_injected_only_for_call_tool() { + let session_id = "test-session-id"; + let tool_call_request_id = "tool-request-1"; + + let call_request = inject_session_context_into_request( + call_tool_request(Extensions::new()), + Some(session_id), + None, + Some(tool_call_request_id), + ); + let call_meta = request_extensions(&call_request) + .and_then(|extensions| extensions.get::()) + .expect("call request should have meta"); + assert_eq!( + call_meta.0.get(TOOL_CALL_REQUEST_ID_HEADER), + Some(&Value::String(tool_call_request_id.to_string())) + ); + + let tools_request = inject_session_context_into_request( + list_tools_request(Extensions::new()), + Some(session_id), + None, + Some(tool_call_request_id), + ); + let tools_meta = request_extensions(&tools_request) + .and_then(|extensions| extensions.get::()) + .expect("list tools request should have meta"); + assert!(!tools_meta.0.contains_key(TOOL_CALL_REQUEST_ID_HEADER)); + } + #[test] fn test_client_info_advertises_mcp_apps_ui_extension() { let client = new_client(GoosePlatform::GooseDesktop); diff --git a/crates/goose/src/agents/platform_extensions/code_execution.rs b/crates/goose/src/agents/platform_extensions/code_execution.rs index 935b51f2f606..6fcf612cab90 100644 --- a/crates/goose/src/agents/platform_extensions/code_execution.rs +++ b/crates/goose/src/agents/platform_extensions/code_execution.rs @@ -143,7 +143,7 @@ impl CodeExecutionClient { /// Build a PctxRegistry with all tool callbacks registered fn build_callback_registry( &self, - session_id: &str, + ctx: &ToolCallContext, code_mode: &CodeMode, ) -> Result { let manager = self @@ -163,7 +163,7 @@ impl CodeExecutionClient { .unwrap_or_default(), &cfg.name ); - let callback = create_tool_callback(session_id.to_string(), full_name, manager.clone()); + let callback = create_tool_callback(ctx.clone(), full_name, manager.clone()); registry .add_callback(&cfg.id(), callback) .map_err(|e| format!("Failed to register callback: {e}"))?; @@ -236,7 +236,7 @@ impl CodeExecutionClient { /// Handle the execute typescript tool call async fn handle_execute_typescript( &self, - session_id: &str, + ctx: &ToolCallContext, arguments: Option, ) -> Result, String> { let args: ExecuteWithToolGraph = arguments @@ -245,8 +245,9 @@ impl CodeExecutionClient { .map_err(|e| format!("Failed to parse arguments: {e}"))? .ok_or("Missing arguments for execute_typescript")?; + let session_id = &ctx.session_id; let code_mode = self.get_code_mode(session_id).await?; - let registry = self.build_callback_registry(session_id, &code_mode)?; + let registry = self.build_callback_registry(ctx, &code_mode)?; let code = args.input.code.clone(); let disclosure = self.disclosure; @@ -273,12 +274,12 @@ impl CodeExecutionClient { } fn create_tool_callback( - session_id: String, + ctx: ToolCallContext, full_name: String, manager: Arc, ) -> CallbackFn { Arc::new(move |args: Option| { - let session_id = session_id.clone(); + let ctx = ctx.clone(); let full_name = full_name.clone(); let manager = manager.clone(); Box::pin(async move { @@ -289,7 +290,6 @@ fn create_tool_callback( } params }; - let ctx = crate::agents::ToolCallContext::new(session_id, None, None); match manager .dispatch_tool_call(&ctx, tool_call, CancellationToken::new()) .await @@ -457,7 +457,7 @@ impl McpClientTrait for CodeExecutionClient { .await } "execute_bash" => self.handle_execute_bash(session_id, arguments).await, - "execute_typescript" => self.handle_execute_typescript(session_id, arguments).await, + "execute_typescript" => self.handle_execute_typescript(ctx, arguments).await, _ => Err(format!("Unknown tool: {name}")), }; diff --git a/crates/goose/src/agents/tool_execution.rs b/crates/goose/src/agents/tool_execution.rs index 603f5d57226c..b5209d622cf7 100644 --- a/crates/goose/src/agents/tool_execution.rs +++ b/crates/goose/src/agents/tool_execution.rs @@ -9,11 +9,13 @@ use tokio_util::sync::CancellationToken; use std::path::PathBuf; use crate::config::permission::PermissionLevel; +use crate::conversation::message::Message; use crate::mcp_utils::ToolResult; use crate::permission::Permission; use rmcp::model::{Content, ServerNotification}; /// Context passed through the tool call dispatch chain. +#[derive(Clone)] pub struct ToolCallContext { pub session_id: String, pub working_dir: Option, @@ -43,6 +45,7 @@ impl ToolCallContext { pub struct ToolCallResult { pub result: Box> + Send + Unpin>, pub notification_stream: Option + Send + Unpin>>, + pub action_required_stream: Option + Send + Unpin>>, } impl From> for ToolCallResult { @@ -50,13 +53,14 @@ impl From> for ToolCallResult { Self { result: Box::new(futures::future::ready(result)), notification_stream: None, + action_required_stream: None, } } } use super::agent::{tool_stream, ToolStream}; use crate::agents::Agent; -use crate::conversation::message::{Message, ToolRequest}; +use crate::conversation::message::ToolRequest; use crate::session::Session; use crate::tool_inspection::get_security_finding_id_from_results; @@ -133,9 +137,11 @@ impl Agent { tool_futures.push((req_id, match tool_result { Ok(result) => tool_stream( result.notification_stream.unwrap_or_else(|| Box::new(stream::empty())), + result.action_required_stream.unwrap_or_else(|| Box::new(stream::empty())), result.result, ), Err(e) => tool_stream( + Box::new(stream::empty()), Box::new(stream::empty()), futures::future::ready(Err(e)), ), diff --git a/crates/goose/src/session_context.rs b/crates/goose/src/session_context.rs index 7bb7cd117a84..ad4046f692e5 100644 --- a/crates/goose/src/session_context.rs +++ b/crates/goose/src/session_context.rs @@ -1,6 +1,7 @@ use tokio::task_local; pub const SESSION_ID_HEADER: &str = "agent-session-id"; +pub const TOOL_CALL_REQUEST_ID_HEADER: &str = "agent-tool-call-request-id"; pub const WORKING_DIR_HEADER: &str = "agent-working-dir"; task_local! { diff --git a/crates/goose/tests/mcp_replays/github-mcp-serverstdio b/crates/goose/tests/mcp_replays/github-mcp-serverstdio index 33f9525dd441..0b6081134897 100644 --- a/crates/goose/tests/mcp_replays/github-mcp-serverstdio +++ b/crates/goose/tests/mcp_replays/github-mcp-serverstdio @@ -9,6 +9,6 @@ STDIN: {"jsonrpc":"2.0","method":"notifications/initialized"} STDERR: time=2025-12-11T17:58:47.642-05:00 level=INFO msg="session initialized" STDIN: {"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":0}}} STDOUT: {"jsonrpc":"2.0","id":1,"result":{"tools":[{"name":"get_file_contents","description":"Get file contents from GitHub","inputSchema":{"type":"object","properties":{"owner":{"type":"string"},"repo":{"type":"string"},"path":{"type":"string"},"sha":{"type":"string"}},"required":["owner","repo","path"]}}]}} -STDIN: {"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":1},"name":"get_file_contents","arguments":{"owner":"block","path":"README.md","repo":"goose","sha":"ab62b863c1666232a67048b6c4e10007a2a5b83c"}}} +STDIN: {"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","agent-tool-call-request-id":"test-id","progressToken":1},"name":"get_file_contents","arguments":{"owner":"block","path":"README.md","repo":"goose","sha":"ab62b863c1666232a67048b6c4e10007a2a5b83c"}}} STDOUT: {"jsonrpc":"2.0","id":2,"result":{"content":[{"type":"text","text":"successfully downloaded text file (SHA: de9bdde7f260549bf3a083651842f30ab29cf4e9)"},{"type":"resource","resource":{"uri":"repo://block/goose/sha/ab62b863c1666232a67048b6c4e10007a2a5b83c/contents/README.md","mimeType":"text/plain; charset=utf-8","text":"\u003cdiv align=\"center\"\u003e\n\n# goose\n\n_a local, extensible, open source AI agent that automates engineering tasks_\n\n\u003cp align=\"center\"\u003e\n \u003ca href=\"https://opensource.org/licenses/Apache-2.0\"\u003e\n \u003cimg src=\"https://img.shields.io/badge/License-Apache_2.0-blue.svg\"\u003e\n \u003c/a\u003e\n \u003ca href=\"https://discord.gg/7GaTvbDwga\"\u003e\n \u003cimg src=\"https://img.shields.io/discord/1287729918100246654?logo=discord\u0026logoColor=white\u0026label=Join+Us\u0026color=blueviolet\" alt=\"Discord\"\u003e\n \u003c/a\u003e\n \u003ca href=\"https://github.com/block/goose/actions/workflows/ci.yml\"\u003e\n \u003cimg src=\"https://img.shields.io/github/actions/workflow/status/block/goose/ci.yml?branch=main\" alt=\"CI\"\u003e\n \u003c/a\u003e\n\u003c/p\u003e\n\u003c/div\u003e\n\ngoose is your on-machine AI agent, capable of automating complex development tasks from start to finish. More than just code suggestions, goose can build entire projects from scratch, write and execute code, debug failures, orchestrate workflows, and interact with external APIs - _autonomously_.\n\nWhether you're prototyping an idea, refining existing code, or managing intricate engineering pipelines, goose adapts to your workflow and executes tasks with precision.\n\nDesigned for maximum flexibility, goose works with any LLM and supports multi-model configuration to optimize performance and cost, seamlessly integrates with MCP servers, and is available as both a desktop app as well as CLI - making it the ultimate AI assistant for developers who want to move faster and focus on innovation.\n\n[![Watch the video](https://github.com/user-attachments/assets/ddc71240-3928-41b5-8210-626dfb28af7a)](https://youtu.be/D-DpDunrbpo)\n\n# Quick Links\n- [Quickstart](https://goose-docs.ai/docs/quickstart)\n- [Installation](https://goose-docs.ai/docs/getting-started/installation)\n- [Tutorials](https://goose-docs.ai/docs/category/tutorials)\n- [Documentation](https://goose-docs.ai/docs/category/getting-started)\n\n\n# a little goose humor 🦢\n\n\u003e Why did the developer choose goose as their AI agent?\n\u003e \n\u003e Because it always helps them \"migrate\" their code to production! 🚀\n\n# goose around with us\n- [Discord](https://discord.gg/block-opensource)\n- [YouTube](https://www.youtube.com/@goose-oss)\n- [LinkedIn](https://www.linkedin.com/company/goose-oss)\n- [Twitter/X](https://x.com/goose_oss)\n- [Bluesky](https://bsky.app/profile/opensource.block.xyz)\n- [Nostr](https://njump.me/opensource@block.xyz)\n"}}]}} STDERR: time=2025-12-11T17:58:48.133-05:00 level=INFO msg="server session disconnected" session_id="" diff --git a/crates/goose/tests/mcp_replays/npx-y@modelcontextprotocol_server-everything@2026.1.14 b/crates/goose/tests/mcp_replays/npx-y@modelcontextprotocol_server-everything@2026.1.14 index c4aaa100a171..21739497a637 100644 --- a/crates/goose/tests/mcp_replays/npx-y@modelcontextprotocol_server-everything@2026.1.14 +++ b/crates/goose/tests/mcp_replays/npx-y@modelcontextprotocol_server-everything@2026.1.14 @@ -6,20 +6,20 @@ STDOUT: {"method":"notifications/tools/list_changed","jsonrpc":"2.0"} STDOUT: {"method":"notifications/tools/list_changed","jsonrpc":"2.0"} STDIN: {"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":0}}} STDOUT: {"jsonrpc":"2.0","id":1,"result":{"tools":[{"name":"echo","description":"Echo a message","inputSchema":{"type":"object","properties":{"message":{"type":"string"}},"required":["message"]}},{"name":"get-sum","description":"Get the sum of two numbers","inputSchema":{"type":"object","properties":{"a":{"type":"number"},"b":{"type":"number"}},"required":["a","b"]}},{"name":"trigger-long-running-operation","description":"Trigger a long-running operation","inputSchema":{"type":"object","properties":{"duration":{"type":"number"},"steps":{"type":"number"}},"required":["duration","steps"]}},{"name":"get-structured-content","description":"Get structured content","inputSchema":{"type":"object","properties":{"location":{"type":"string"}},"required":["location"]}},{"name":"trigger-sampling-request","description":"Trigger a sampling request","inputSchema":{"type":"object","properties":{"prompt":{"type":"string"},"maxTokens":{"type":"number"}},"required":["prompt","maxTokens"]}}]}} -STDIN: {"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":1},"name":"echo","arguments":{"message":"Hello, world!"}}} +STDIN: {"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","agent-tool-call-request-id":"test-id","progressToken":1},"name":"echo","arguments":{"message":"Hello, world!"}}} STDOUT: {"result":{"content":[{"type":"text","text":"Echo: Hello, world!"}]},"jsonrpc":"2.0","id":2} -STDIN: {"jsonrpc":"2.0","id":3,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":2},"name":"get-sum","arguments":{"a":1,"b":2}}} +STDIN: {"jsonrpc":"2.0","id":3,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","agent-tool-call-request-id":"test-id","progressToken":2},"name":"get-sum","arguments":{"a":1,"b":2}}} STDOUT: {"result":{"content":[{"type":"text","text":"The sum of 1 and 2 is 3."}]},"jsonrpc":"2.0","id":3} -STDIN: {"jsonrpc":"2.0","id":4,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":3},"name":"trigger-long-running-operation","arguments":{"duration":1,"steps":5}}} +STDIN: {"jsonrpc":"2.0","id":4,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","agent-tool-call-request-id":"test-id","progressToken":3},"name":"trigger-long-running-operation","arguments":{"duration":1,"steps":5}}} STDOUT: {"method":"notifications/progress","params":{"progress":1,"total":5,"progressToken":3},"jsonrpc":"2.0"} STDOUT: {"method":"notifications/progress","params":{"progress":2,"total":5,"progressToken":3},"jsonrpc":"2.0"} STDOUT: {"method":"notifications/progress","params":{"progress":3,"total":5,"progressToken":3},"jsonrpc":"2.0"} STDOUT: {"method":"notifications/progress","params":{"progress":4,"total":5,"progressToken":3},"jsonrpc":"2.0"} STDOUT: {"method":"notifications/progress","params":{"progress":5,"total":5,"progressToken":3},"jsonrpc":"2.0"} STDOUT: {"result":{"content":[{"type":"text","text":"Long running operation completed. Duration: 1 seconds, Steps: 5."}]},"jsonrpc":"2.0","id":4} -STDIN: {"jsonrpc":"2.0","id":5,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":4},"name":"get-structured-content","arguments":{"location":"New York"}}} +STDIN: {"jsonrpc":"2.0","id":5,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","agent-tool-call-request-id":"test-id","progressToken":4},"name":"get-structured-content","arguments":{"location":"New York"}}} STDOUT: {"result":{"content":[{"type":"text","text":"{\"temperature\":33,\"conditions\":\"Cloudy\",\"humidity\":82}"}],"structuredContent":{"temperature":33,"conditions":"Cloudy","humidity":82}},"jsonrpc":"2.0","id":5} -STDIN: {"jsonrpc":"2.0","id":6,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":5},"name":"trigger-sampling-request","arguments":{"maxTokens":100,"prompt":"Please provide a quote from The Great Gatsby"}}} +STDIN: {"jsonrpc":"2.0","id":6,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","agent-tool-call-request-id":"test-id","progressToken":5},"name":"trigger-sampling-request","arguments":{"maxTokens":100,"prompt":"Please provide a quote from The Great Gatsby"}}} STDOUT: {"method":"sampling/createMessage","params":{"messages":[{"role":"user","content":{"type":"text","text":"Resource trigger-sampling-request context: Please provide a quote from The Great Gatsby"}}],"systemPrompt":"You are a helpful test server.","maxTokens":100,"temperature":0.7},"jsonrpc":"2.0","id":0} STDIN: {"jsonrpc":"2.0","id":0,"result":{"model":"mock","stopReason":"endTurn","role":"assistant","content":{"type":"text","text":"\"So we beat on, boats against the current, borne back ceaselessly into the past.\" — F. Scott Fitzgerald, The Great Gatsby (1925)"}}} STDOUT: {"result":{"content":[{"type":"text","text":"LLM sampling result: \n{\n \"model\": \"mock\",\n \"stopReason\": \"endTurn\",\n \"role\": \"assistant\",\n \"content\": {\n \"type\": \"text\",\n \"text\": \"\\\"So we beat on, boats against the current, borne back ceaselessly into the past.\\\" — F. Scott Fitzgerald, The Great Gatsby (1925)\"\n }\n}"}]},"jsonrpc":"2.0","id":6} diff --git a/crates/goose/tests/mcp_replays/uvrun--withfastmcp==2.14.4fastmcpruntests_fastmcp_test_server.py b/crates/goose/tests/mcp_replays/uvrun--withfastmcp==2.14.4fastmcpruntests_fastmcp_test_server.py index 4494d832312b..9b289d031ad1 100644 --- a/crates/goose/tests/mcp_replays/uvrun--withfastmcp==2.14.4fastmcpruntests_fastmcp_test_server.py +++ b/crates/goose/tests/mcp_replays/uvrun--withfastmcp==2.14.4fastmcpruntests_fastmcp_test_server.py @@ -27,5 +27,5 @@ STDIN: {"jsonrpc":"2.0","method":"notifications/initialized"} STDIN: {"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":0}}} STDOUT: {"jsonrpc":"2.0","id":1,"result":{"tools":[{"name":"divide","description":"Divide two numbers","inputSchema":{"type":"object","properties":{"dividend":{"type":"number"},"divisor":{"type":"number"}},"required":["dividend","divisor"]}}]}} -STDIN: {"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":1},"name":"divide","arguments":{"dividend":10,"divisor":2}}} +STDIN: {"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","agent-tool-call-request-id":"test-id","progressToken":1},"name":"divide","arguments":{"dividend":10,"divisor":2}}} STDOUT: {"jsonrpc":"2.0","id":2,"result":{"content":[{"type":"text","text":"5.0"}],"structuredContent":{"result":5.0},"isError":false}} diff --git a/crates/goose/tests/mcp_replays/uvxmcp-server-fetch b/crates/goose/tests/mcp_replays/uvxmcp-server-fetch index 859a737e53d9..9531c8c2a8c3 100644 --- a/crates/goose/tests/mcp_replays/uvxmcp-server-fetch +++ b/crates/goose/tests/mcp_replays/uvxmcp-server-fetch @@ -3,5 +3,5 @@ STDOUT: {"jsonrpc":"2.0","id":0,"result":{"protocolVersion":"2025-03-26","capabi STDIN: {"jsonrpc":"2.0","method":"notifications/initialized"} STDIN: {"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":0}}} STDOUT: {"jsonrpc":"2.0","id":1,"result":{"tools":[{"name":"fetch","description":"Fetch a URL","inputSchema":{"type":"object","properties":{"url":{"type":"string"}},"required":["url"]}}]}} -STDIN: {"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","progressToken":1},"name":"fetch","arguments":{"url":"https://example.com"}}} +STDIN: {"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"_meta":{"agent-session-id":"test-session-id","agent-tool-call-request-id":"test-id","progressToken":1},"name":"fetch","arguments":{"url":"https://example.com"}}} STDOUT: {"jsonrpc":"2.0","id":2,"result":{"content":[{"type":"text","text":"Contents of https://example.com/:\nThis domain is for use in documentation examples without needing permission. Avoid use in operations.\n\n[Learn more](https://iana.org/domains/example)"}],"isError":false}} From 8e7d3f56558d706a6d7fc7cc39ffb8f2507c417b Mon Sep 17 00:00:00 2001 From: Douwe Osinga Date: Tue, 23 Jun 2026 21:32:56 -0400 Subject: [PATCH 10/12] Migrate diagnostics to JSON report (#9964) Co-authored-by: Douwe M Osinga --- Cargo.lock | 1 - crates/goose-cli/src/cli.rs | 2 - crates/goose-cli/src/commands/session.rs | 29 +- crates/goose-cli/src/commands/term.rs | 12 +- crates/goose-sdk-types/src/custom_requests.rs | 25 ++ crates/goose-server/src/openapi.rs | 15 +- crates/goose-server/src/routes/status.rs | 51 ++- crates/goose/Cargo.toml | 1 - crates/goose/acp-meta.json | 5 + crates/goose/acp-schema.json | 52 +++ crates/goose/src/acp/server.rs | 1 + .../goose/src/acp/server/custom_dispatch.rs | 8 + crates/goose/src/acp/server/diagnostics.rs | 20 ++ crates/goose/src/session/diagnostics.rs | 313 ++++++++++++++---- crates/goose/src/session/mod.rs | 4 +- crates/goose/src/session/session_manager.rs | 120 +++++++ .../docs/guides/goose-cli-commands.md | 10 +- .../diagnostics-and-reporting.md | 36 +- scripts/diagnostics-viewer.py | 130 ++++++-- ui/desktop/openapi.json | 214 +++++++++++- ui/desktop/src/acp/diagnostics.ts | 16 + ui/desktop/src/api/index.ts | 2 +- ui/desktop/src/api/types.gen.ts | 61 +++- ui/desktop/src/components/ui/Diagnostics.tsx | 24 +- ui/desktop/src/i18n/messages/en.json | 4 +- ui/sdk/src/generated/client.gen.ts | 15 + ui/sdk/src/generated/index.ts | 7 +- ui/sdk/src/generated/types.gen.ts | 15 +- ui/sdk/src/generated/zod.gen.ts | 13 + 29 files changed, 1029 insertions(+), 177 deletions(-) create mode 100644 crates/goose/src/acp/server/diagnostics.rs create mode 100644 ui/desktop/src/acp/diagnostics.ts diff --git a/Cargo.lock b/Cargo.lock index 122aadde02f8..899297ffaa9c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4886,7 +4886,6 @@ dependencies = [ "which 8.0.3", "winapi", "wiremock", - "zip 8.6.0", ] [[package]] diff --git a/crates/goose-cli/src/cli.rs b/crates/goose-cli/src/cli.rs index 3694127a4e20..267c0b352258 100644 --- a/crates/goose-cli/src/cli.rs +++ b/crates/goose-cli/src/cli.rs @@ -580,11 +580,9 @@ enum SessionCommand { }, #[command(name = "diagnostics")] Diagnostics { - /// Session identifier for generating diagnostics #[command(flatten)] identifier: Option, - /// Output path for the diagnostics zip file (optional, defaults to current directory) #[arg(short = 'o', long)] output: Option, }, diff --git a/crates/goose-cli/src/commands/session.rs b/crates/goose-cli/src/commands/session.rs index 1a91a654bdd1..15d5908888d8 100644 --- a/crates/goose-cli/src/commands/session.rs +++ b/crates/goose-cli/src/commands/session.rs @@ -7,7 +7,9 @@ use etcetera::home_dir; use goose::config::Config; #[cfg(feature = "nostr")] use goose::session::nostr_share; -use goose::session::{generate_diagnostics, Session, SessionManager, SessionType}; +use goose::session::{ + generate_diagnostics, DiagnosticsLevel, Session, SessionManager, SessionType, +}; use goose::utils::safe_truncate; use regex::Regex; use std::fs; @@ -322,24 +324,27 @@ pub async fn handle_session_import(input: String, nostr: bool) -> Result<()> { pub async fn handle_diagnostics(session_id: &str, output_path: Option) -> Result<()> { println!( - "Generating diagnostics bundle for session '{}'...", + "Generating diagnostics report for session '{}'...", session_id ); let session_manager = SessionManager::instance(); - let diagnostics_data = generate_diagnostics(&session_manager, session_id) - .await - .with_context(|| { - format!( - "Failed to write to generate diagnostics bundle for session '{}'", - session_id - ) - })?; + let diagnostics_report = + generate_diagnostics(&session_manager, session_id, DiagnosticsLevel::Full) + .await + .with_context(|| { + format!( + "Failed to generate diagnostics report for session '{}'", + session_id + ) + })?; + let diagnostics_data = serde_json::to_vec_pretty(&diagnostics_report) + .context("Failed to serialize diagnostics report")?; let output_file = if let Some(path) = output_path { path.clone() } else { - PathBuf::from(format!("diagnostics_{}.zip", session_id)) + PathBuf::from(format!("diagnostics_{}.json", session_id)) }; let mut file = fs::File::create(&output_file).context(format!( @@ -350,7 +355,7 @@ pub async fn handle_diagnostics(session_id: &str, output_path: Option) file.write_all(&diagnostics_data) .context("Failed to write diagnostics data")?; - println!("Diagnostics bundle saved to: {}", output_file.display()); + println!("Diagnostics report saved to: {}", output_file.display()); Ok(()) } diff --git a/crates/goose-cli/src/commands/term.rs b/crates/goose-cli/src/commands/term.rs index 26d2a3652d06..9c80fc1e0e3e 100644 --- a/crates/goose-cli/src/commands/term.rs +++ b/crates/goose-cli/src/commands/term.rs @@ -280,9 +280,15 @@ pub async fn handle_term_run(prompt: Vec) -> Result<()> { }; if let Some(oldest_user) = user_messages_after_last_assistant.last() { - session_manager - .truncate_conversation(&session_id, oldest_user.created) - .await?; + if let Some(message_id) = oldest_user.id.as_deref() { + session_manager + .truncate_conversation_from_message(&session_id, message_id) + .await?; + } else { + session_manager + .truncate_conversation(&session_id, oldest_user.created) + .await?; + } } let prompt_with_context = if user_messages_after_last_assistant.is_empty() { diff --git a/crates/goose-sdk-types/src/custom_requests.rs b/crates/goose-sdk-types/src/custom_requests.rs index 48e1bc10b21e..89061ba33e1f 100644 --- a/crates/goose-sdk-types/src/custom_requests.rs +++ b/crates/goose-sdk-types/src/custom_requests.rs @@ -165,6 +165,31 @@ pub struct SteerSessionResponse { pub message_id: String, } +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request( + method = "_goose/unstable/diagnostics/get", + response = DiagnosticsGetResponse +)] +#[serde(rename_all = "camelCase")] +pub struct DiagnosticsGetRequest { + pub session_id: String, + #[serde(default)] + pub level: DiagnosticsReportLevel, +} + +#[derive(Debug, Default, Clone, Copy, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "snake_case")] +pub enum DiagnosticsReportLevel { + #[default] + Summary, + Full, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +pub struct DiagnosticsGetResponse { + pub report: serde_json::Value, +} + /// Delete a session. #[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] #[request(method = "session/delete", response = EmptyResponse)] diff --git a/crates/goose-server/src/openapi.rs b/crates/goose-server/src/openapi.rs index 823feb29d5c3..9e2ef4d864dd 100644 --- a/crates/goose-server/src/openapi.rs +++ b/crates/goose-server/src/openapi.rs @@ -7,7 +7,11 @@ use goose::conversation::token_usage::Usage; use goose::conversation::Conversation; use goose::download_manager::{DownloadProgress, DownloadStatus}; use goose::providers::base::{ConfigKey, ModelInfo, ProviderMetadata, ProviderType}; -use goose::session::{Session, SessionType, SystemInfo}; +use goose::session::{ + DiagnosticsConfig, DiagnosticsError, DiagnosticsExtensions, DiagnosticsLevel, DiagnosticsLogs, + DiagnosticsPrompt, DiagnosticsReport, DiagnosticsScheduledRecipe, DiagnosticsTextFile, Session, + SessionType, SystemInfo, +}; use goose_providers::model::ModelConfig; use goose_providers::permission::Permission; use goose_providers::permission::PrincipalType; @@ -583,6 +587,15 @@ derive_utoipa!(IconTheme as IconThemeSchema); goose_providers::goose_mode::GooseMode, SessionType, SystemInfo, + DiagnosticsConfig, + DiagnosticsError, + DiagnosticsExtensions, + DiagnosticsLevel, + DiagnosticsLogs, + DiagnosticsPrompt, + DiagnosticsReport, + DiagnosticsScheduledRecipe, + DiagnosticsTextFile, Conversation, IconSchema, IconThemeSchema, diff --git a/crates/goose-server/src/routes/status.rs b/crates/goose-server/src/routes/status.rs index edc6c71c6300..ec86e6316b1d 100644 --- a/crates/goose-server/src/routes/status.rs +++ b/crates/goose-server/src/routes/status.rs @@ -1,9 +1,9 @@ -use axum::body::Body; -use axum::extract::State; -use axum::http::HeaderValue; -use axum::response::IntoResponse; -use axum::{extract::Path, http::StatusCode, routing::get, Json, Router}; -use goose::session::{generate_diagnostics, get_system_info, SystemInfo}; +use axum::extract::{Path, Query, State}; +use axum::{http::StatusCode, routing::get, Json, Router}; +use goose::session::{ + generate_diagnostics, get_system_info, DiagnosticsLevel, DiagnosticsReport, SystemInfo, +}; +use serde::Deserialize; use std::sync::Arc; use crate::state::AppState; @@ -26,34 +26,33 @@ async fn system_info() -> Json { Json(get_system_info()) } +#[derive(Debug, Default, Deserialize, utoipa::IntoParams)] +struct DiagnosticsQuery { + level: Option, +} + #[utoipa::path(get, path = "/diagnostics/{session_id}", + params( + DiagnosticsQuery, + ), responses( - (status = 200, description = "Diagnostics zip file", content_type = "application/zip", body = Vec), + (status = 200, description = "Diagnostics report", body = DiagnosticsReport), (status = 500, description = "Failed to generate diagnostics"), ) )] async fn diagnostics( State(state): State>, Path(session_id): Path, -) -> impl IntoResponse { - match generate_diagnostics(state.session_manager(), &session_id).await { - Ok(zip_data) => { - let filename = format!("attachment; filename=\"diagnostics_{}.zip\"", session_id); - let headers = [ - ( - http::header::CONTENT_TYPE, - HeaderValue::from_static("application/zip"), - ), - ( - http::header::CONTENT_DISPOSITION, - HeaderValue::from_str(&filename).map_err(|_e| StatusCode::BAD_REQUEST)?, - ), - ]; - - Ok((headers, Body::from(zip_data))) - } - Err(_) => Err(StatusCode::INTERNAL_SERVER_ERROR), - } + Query(query): Query, +) -> Result, StatusCode> { + generate_diagnostics( + state.session_manager(), + &session_id, + query.level.unwrap_or(DiagnosticsLevel::Full), + ) + .await + .map(Json) + .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR) } pub fn routes(state: Arc) -> Router { Router::new() diff --git a/crates/goose/Cargo.toml b/crates/goose/Cargo.toml index 294bcaee95f8..9a52ef60f3e3 100644 --- a/crates/goose/Cargo.toml +++ b/crates/goose/Cargo.toml @@ -168,7 +168,6 @@ byteorder = { version = "1.5", default-features = false, features = ["std"], opt tokenizers = { version = "0.23", default-features = false, features = ["onig"], optional = true } symphonia = { version = "0.5", default-features = false, features = ["aac", "adpcm", "alac", "isomp4", "mkv", "mp3", "pcm", "vorbis", "wav"], optional = true } rubato = { version = "0.16", default-features = false, optional = true } -zip = { workspace = true } sys-info = { version = "0.9", default-features = false } llama-cpp-2 = { workspace = true, optional = true } diff --git a/crates/goose/acp-meta.json b/crates/goose/acp-meta.json index eef554cdf90f..f0efcda0fb45 100644 --- a/crates/goose/acp-meta.json +++ b/crates/goose/acp-meta.json @@ -40,6 +40,11 @@ "requestType": "SteerSessionRequest_unstable", "responseType": "SteerSessionResponse_unstable" }, + { + "method": "_goose/unstable/diagnostics/get", + "requestType": "DiagnosticsGetRequest_unstable", + "responseType": "DiagnosticsGetResponse_unstable" + }, { "method": "session/delete", "requestType": "DeleteSessionRequest", diff --git a/crates/goose/acp-schema.json b/crates/goose/acp-schema.json index 1cf4fb89f017..9124a9a363ef 100644 --- a/crates/goose/acp-schema.json +++ b/crates/goose/acp-schema.json @@ -973,6 +973,41 @@ "x-side": "agent", "x-method": "_goose/unstable/session/steer" }, + "DiagnosticsGetRequest_unstable": { + "type": "object", + "properties": { + "sessionId": { + "type": "string" + }, + "level": { + "$ref": "#/$defs/DiagnosticsReportLevel", + "default": "summary" + } + }, + "required": [ + "sessionId" + ], + "x-side": "agent", + "x-method": "_goose/unstable/diagnostics/get" + }, + "DiagnosticsReportLevel": { + "type": "string", + "enum": [ + "summary", + "full" + ] + }, + "DiagnosticsGetResponse_unstable": { + "type": "object", + "properties": { + "report": {} + }, + "required": [ + "report" + ], + "x-side": "agent", + "x-method": "_goose/unstable/diagnostics/get" + }, "DeleteSessionRequest": { "type": "object", "properties": { @@ -4613,6 +4648,15 @@ "description": "Params for _goose/unstable/session/steer", "title": "SteerSessionRequest_unstable" }, + { + "allOf": [ + { + "$ref": "#/$defs/DiagnosticsGetRequest_unstable" + } + ], + "description": "Params for _goose/unstable/diagnostics/get", + "title": "DiagnosticsGetRequest_unstable" + }, { "allOf": [ { @@ -5250,6 +5294,14 @@ ], "title": "SteerSessionResponse_unstable" }, + { + "allOf": [ + { + "$ref": "#/$defs/DiagnosticsGetResponse_unstable" + } + ], + "title": "DiagnosticsGetResponse_unstable" + }, { "allOf": [ { diff --git a/crates/goose/src/acp/server.rs b/crates/goose/src/acp/server.rs index ae2d1e1acdbb..c8e9f0fb5161 100644 --- a/crates/goose/src/acp/server.rs +++ b/crates/goose/src/acp/server.rs @@ -84,6 +84,7 @@ mod agent_requests; pub use agent_requests::agent_request_schemas; mod config; mod custom_dispatch; +mod diagnostics; mod dictation; mod dispatch; mod elicitation; diff --git a/crates/goose/src/acp/server/custom_dispatch.rs b/crates/goose/src/acp/server/custom_dispatch.rs index 1652023637ce..7b10cc3178ca 100644 --- a/crates/goose/src/acp/server/custom_dispatch.rs +++ b/crates/goose/src/acp/server/custom_dispatch.rs @@ -82,6 +82,14 @@ impl GooseAcpAgent { self.on_steer_session(req).await } + #[custom_method(DiagnosticsGetRequest)] + async fn dispatch_get_diagnostics( + &self, + req: DiagnosticsGetRequest, + ) -> Result { + self.on_get_diagnostics(req).await + } + #[custom_method(DeleteSessionRequest)] async fn dispatch_delete_session( &self, diff --git a/crates/goose/src/acp/server/diagnostics.rs b/crates/goose/src/acp/server/diagnostics.rs new file mode 100644 index 000000000000..954aeed7aa05 --- /dev/null +++ b/crates/goose/src/acp/server/diagnostics.rs @@ -0,0 +1,20 @@ +use super::*; +use crate::session::{generate_diagnostics, DiagnosticsLevel}; + +impl GooseAcpAgent { + pub(super) async fn on_get_diagnostics( + &self, + req: DiagnosticsGetRequest, + ) -> Result { + let level = match req.level { + DiagnosticsReportLevel::Summary => DiagnosticsLevel::Summary, + DiagnosticsReportLevel::Full => DiagnosticsLevel::Full, + }; + let report = generate_diagnostics(&self.session_manager, &req.session_id, level) + .await + .internal_err()?; + let report = serde_json::to_value(report).internal_err()?; + + Ok(DiagnosticsGetResponse { report }) + } +} diff --git a/crates/goose/src/session/diagnostics.rs b/crates/goose/src/session/diagnostics.rs index c9e580c7f6da..bbc84a9022df 100644 --- a/crates/goose/src/session/diagnostics.rs +++ b/crates/goose/src/session/diagnostics.rs @@ -6,14 +6,22 @@ use crate::providers::utils::LOGS_TO_KEEP; use crate::session::SessionManager; use serde::{Deserialize, Serialize}; use std::fs; -use std::io::Cursor; -use std::io::Write; use std::path::PathBuf; use utoipa::ToSchema; -use zip::write::SimpleFileOptions; -use zip::ZipWriter; -#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)] +const SERVER_LOG_TAIL_LINES: usize = 400; +const LLM_LOG_MAX_BYTES: usize = 2 * 1024 * 1024; +const CONFIG_MAX_BYTES: usize = 256 * 1024; + +#[derive(Debug, Clone, Copy, Default, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] +#[serde(rename_all = "snake_case")] +pub enum DiagnosticsLevel { + #[default] + Summary, + Full, +} + +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] pub struct SystemInfo { pub app_version: String, pub os: String, @@ -24,6 +32,73 @@ pub struct SystemInfo { pub enabled_extensions: Vec, } +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct DiagnosticsConfig { + pub config_path: String, + pub config_yaml: Option, + pub truncated: bool, +} + +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct DiagnosticsExtensions { + pub enabled: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct DiagnosticsTextFile { + pub path: String, + pub content: String, + pub truncated: bool, +} + +#[derive(Debug, Clone, Default, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct DiagnosticsLogs { + pub server: Option, + pub llm: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct DiagnosticsPrompt { + pub name: String, + pub content: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct DiagnosticsScheduledRecipe { + pub path: String, + pub content: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct DiagnosticsError { + pub path: Option, + pub message: String, +} + +#[derive(Debug, Clone, Serialize, Deserialize, ToSchema, schemars::JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct DiagnosticsReport { + pub schema_version: u32, + pub generated_at: String, + pub level: DiagnosticsLevel, + pub system: SystemInfo, + pub config: Option, + pub extensions: DiagnosticsExtensions, + pub session: Option, + pub logs: DiagnosticsLogs, + pub prompts: Vec, + pub schedule: Option, + pub scheduled_recipes: Vec, + pub errors: Vec, +} + impl SystemInfo { pub fn collect() -> Self { let config = Config::global(); @@ -86,6 +161,68 @@ pub fn latest_llm_log_path() -> Option { path.exists().then_some(path) } +fn recent_llm_log_paths() -> Vec { + let logs_dir = Paths::in_state_dir("logs"); + let paths: Vec<_> = fs::read_dir(logs_dir) + .ok() + .into_iter() + .flatten() + .filter_map(|entry| entry.ok().map(|entry| entry.path())) + .filter(|path| { + path.file_name() + .and_then(|name| name.to_str()) + .is_some_and(|name| name.starts_with("llm_request.") && name.ends_with(".jsonl")) + }) + .collect(); + + let (mut numbered, mut temp): (Vec<_>, Vec<_>) = paths + .into_iter() + .partition(|path| llm_log_index(path).is_some()); + + numbered.sort_by_key(|path| llm_log_index(path).unwrap_or(usize::MAX)); + temp.sort_by(|left, right| { + llm_log_modified(right) + .cmp(&llm_log_modified(left)) + .then_with(|| llm_log_name(left).cmp(&llm_log_name(right))) + }); + + if temp.is_empty() || numbered.len() < LOGS_TO_KEEP { + numbered.extend(temp); + numbered.truncate(LOGS_TO_KEEP); + numbered + } else { + let temp_slots = 1; + let numbered_slots = LOGS_TO_KEEP.saturating_sub(temp_slots); + temp.truncate(temp_slots); + numbered.truncate(numbered_slots); + temp.extend(numbered); + temp + } +} + +fn llm_log_index(path: &std::path::Path) -> Option { + let name = path + .file_name() + .and_then(|name| name.to_str()) + .unwrap_or_default(); + name.strip_prefix("llm_request.") + .and_then(|name| name.strip_suffix(".jsonl")) + .and_then(|name| name.parse::().ok()) +} + +fn llm_log_modified(path: &std::path::Path) -> std::time::SystemTime { + path.metadata() + .and_then(|metadata| metadata.modified()) + .unwrap_or(std::time::SystemTime::UNIX_EPOCH) +} + +fn llm_log_name(path: &std::path::Path) -> String { + path.file_name() + .and_then(|name| name.to_str()) + .unwrap_or_default() + .to_string() +} + pub fn read_tail(path: &std::path::Path, max_lines: usize) -> Option { let content = fs::read_to_string(path).ok()?; let lines: Vec<&str> = content.lines().collect(); @@ -128,6 +265,10 @@ pub fn read_capped(path: &std::path::Path, max_bytes: usize) -> Option { )) } +fn was_truncated(content: &str) -> bool { + content.contains("... (") && content.contains(" bytes omitted) ...") +} + fn latest_entry_by_name(dir: &std::path::Path) -> Option { let mut entries: Vec<_> = fs::read_dir(dir).ok()?.filter_map(|e| e.ok()).collect(); entries.sort_by_key(|e| e.file_name()); @@ -137,81 +278,135 @@ fn latest_entry_by_name(dir: &std::path::Path) -> Option { pub async fn generate_diagnostics( session_manager: &SessionManager, session_id: &str, -) -> anyhow::Result> { - let logs_dir = Paths::in_state_dir("logs"); - let config_dir = Paths::config_dir(); - let config_path = config_dir.join("config.yaml"); + level: DiagnosticsLevel, +) -> anyhow::Result { + let config_path = config_path(); let data_dir = Paths::data_dir(); - let system_info = SystemInfo::collect(); + let is_full = matches!(level, DiagnosticsLevel::Full); + let mut errors = Vec::new(); - let mut buffer = Vec::new(); - { - let mut zip = ZipWriter::new(Cursor::new(&mut buffer)); - let options = - SimpleFileOptions::default().compression_method(zip::CompressionMethod::Deflated); - - let mut log_files: Vec<_> = fs::read_dir(&logs_dir)? - .filter_map(|e| e.ok()) - .filter(|e| e.path().extension().is_some_and(|ext| ext == "jsonl")) - .collect(); - - log_files.sort_by_key(|e| e.metadata().ok().and_then(|m| m.modified().ok())); - - for entry in log_files.iter().rev().take(LOGS_TO_KEEP) { - let path = entry.path(); - let name = path.file_name().unwrap().to_str().unwrap(); - zip.start_file(format!("logs/{}", name), options)?; - zip.write_all(&fs::read(&path)?)?; - } - - if let Some(server_log) = latest_server_log_path() { - if let Ok(content) = fs::read(&server_log) { - let name = server_log.file_name().unwrap().to_str().unwrap(); - zip.start_file(format!("logs/server/{}", name), options)?; - zip.write_all(&content)?; - } - } - + let session = if is_full { let session_data = session_manager.export_session(session_id).await?; - zip.start_file("session.json", options)?; - zip.write_all(session_data.as_bytes())?; + Some(serde_json::from_str(&session_data)?) + } else { + None + }; - if config_path.exists() { - zip.start_file("config.yaml", options)?; - zip.write_all(&fs::read(&config_path)?)?; + let config = if is_full { + let config_yaml = if config_path.exists() { + read_capped(&config_path, CONFIG_MAX_BYTES) + } else { + None + }; + let truncated = config_yaml.as_deref().is_some_and(was_truncated); + Some(DiagnosticsConfig { + config_path: config_path.display().to_string(), + config_yaml, + truncated, + }) + } else { + None + }; + + let logs = if is_full { + DiagnosticsLogs { + server: latest_server_log_path().and_then(|path| { + read_tail(&path, SERVER_LOG_TAIL_LINES).map(|content| DiagnosticsTextFile { + path: path.display().to_string(), + content, + truncated: true, + }) + }), + llm: recent_llm_log_paths() + .into_iter() + .filter_map(|path| { + read_capped(&path, LLM_LOG_MAX_BYTES).map(|content| { + let truncated = was_truncated(&content); + DiagnosticsTextFile { + path: path.display().to_string(), + content, + truncated, + } + }) + }) + .collect(), } + } else { + DiagnosticsLogs::default() + }; - zip.start_file("system.txt", options)?; - zip.write_all(system_info.to_text().as_bytes())?; + let prompts = if is_full { + list_templates() + .into_iter() + .map(|template| DiagnosticsPrompt { + name: template.name, + content: template.user_content.unwrap_or(template.default_content), + }) + .collect() + } else { + Vec::new() + }; + let schedule = if is_full { let schedule_json = data_dir.join("schedule.json"); if schedule_json.exists() { - zip.start_file("schedule.json", options)?; - zip.write_all(&fs::read(&schedule_json)?)?; + fs::read_to_string(&schedule_json).ok().and_then(|content| { + match serde_json::from_str(&content) { + Ok(value) => Some(value), + Err(err) => { + errors.push(DiagnosticsError { + path: Some(schedule_json.display().to_string()), + message: err.to_string(), + }); + None + } + } + }) + } else { + None } + } else { + None + }; + let mut scheduled_recipes = Vec::new(); + if is_full { let scheduled_recipes_dir = data_dir.join("scheduled_recipes"); if scheduled_recipes_dir.exists() && scheduled_recipes_dir.is_dir() { for entry in fs::read_dir(&scheduled_recipes_dir)? { let entry = entry?; let path = entry.path(); if path.is_file() { - let name = path.file_name().unwrap().to_str().unwrap(); - zip.start_file(format!("scheduled_recipes/{}", name), options)?; - zip.write_all(&fs::read(&path)?)?; + match fs::read_to_string(&path) { + Ok(content) => scheduled_recipes.push(DiagnosticsScheduledRecipe { + path: path.display().to_string(), + content, + }), + Err(err) => errors.push(DiagnosticsError { + path: Some(path.display().to_string()), + message: err.to_string(), + }), + } } } } - - for template in list_templates() { - let content = template.user_content.unwrap_or(template.default_content); - zip.start_file(format!("prompts/{}.txt", template.name), options)?; - zip.write_all(content.as_bytes())?; - } - - zip.finish()?; } - Ok(buffer) + Ok(DiagnosticsReport { + schema_version: 1, + generated_at: chrono::Utc::now().to_rfc3339(), + level, + system: system_info.clone(), + config, + extensions: DiagnosticsExtensions { + enabled: system_info.enabled_extensions, + }, + session, + logs, + prompts, + schedule, + scheduled_recipes, + errors, + }) } diff --git a/crates/goose/src/session/mod.rs b/crates/goose/src/session/mod.rs index 951c533bcd3b..c7c5230e8013 100644 --- a/crates/goose/src/session/mod.rs +++ b/crates/goose/src/session/mod.rs @@ -11,7 +11,9 @@ mod session_naming; pub use diagnostics::{ config_path, generate_diagnostics, get_system_info, latest_llm_log_path, - latest_server_log_path, read_capped, read_tail, SystemInfo, + latest_server_log_path, read_capped, read_tail, DiagnosticsConfig, DiagnosticsError, + DiagnosticsExtensions, DiagnosticsLevel, DiagnosticsLogs, DiagnosticsPrompt, DiagnosticsReport, + DiagnosticsScheduledRecipe, DiagnosticsTextFile, SystemInfo, }; pub use extension_data::{EnabledExtensionsState, ExtensionData, ExtensionState, TodoState}; pub use session_manager::{ diff --git a/crates/goose/src/session/session_manager.rs b/crates/goose/src/session/session_manager.rs index 6b7971211657..7b9e3f43de55 100644 --- a/crates/goose/src/session/session_manager.rs +++ b/crates/goose/src/session/session_manager.rs @@ -451,6 +451,16 @@ impl SessionManager { .await } + pub async fn truncate_conversation_from_message( + &self, + session_id: &str, + message_id: &str, + ) -> Result<()> { + self.storage + .truncate_conversation_from_message(session_id, message_id) + .await + } + async fn system_generated_name_update( &self, id: &str, @@ -1948,6 +1958,38 @@ impl SessionStorage { Ok(()) } + async fn truncate_conversation_from_message( + &self, + session_id: &str, + message_id: &str, + ) -> Result<()> { + let pool = self.pool().await?; + let mut tx = pool.begin_with("BEGIN IMMEDIATE").await?; + + let boundary = sqlx::query_as::<_, (i64, i64)>( + "SELECT id, created_timestamp FROM messages WHERE session_id = ? AND message_id = ? ORDER BY created_timestamp, id LIMIT 1", + ) + .bind(session_id) + .bind(message_id) + .fetch_optional(&mut *tx) + .await?; + + if let Some((boundary_id, boundary_timestamp)) = boundary { + sqlx::query( + "DELETE FROM messages WHERE session_id = ? AND (created_timestamp > ? OR (created_timestamp = ? AND id >= ?))", + ) + .bind(session_id) + .bind(boundary_timestamp) + .bind(boundary_timestamp) + .bind(boundary_id) + .execute(&mut *tx) + .await?; + } + + tx.commit().await?; + Ok(()) + } + async fn search_chat_history( &self, query: &str, @@ -2227,12 +2269,90 @@ mod tests { .unwrap(); } + async fn set_message_timestamp( + sm: &SessionManager, + session_id: &str, + message_id: &str, + timestamp: &str, + ) { + let pool = sm.storage().pool().await.unwrap(); + let timestamp = chrono::DateTime::parse_from_rfc3339(timestamp).unwrap(); + let timestamp_string = timestamp.format("%Y-%m-%d %H:%M:%S").to_string(); + + sqlx::query( + "UPDATE messages SET timestamp = ?, created_timestamp = ? WHERE session_id = ? AND message_id = ?", + ) + .bind(×tamp_string) + .bind(timestamp.timestamp()) + .bind(session_id) + .bind(message_id) + .execute(pool) + .await + .unwrap(); + } + async fn add_user_message(sm: &SessionManager, session_id: &str) { sm.add_message(session_id, &Message::user().with_text("hello world")) .await .unwrap(); } + #[tokio::test] + async fn test_truncate_conversation_from_message_keeps_same_second_previous_rows() { + let temp_dir = TempDir::new().unwrap(); + let sm = SessionManager::new(temp_dir.path().to_path_buf()); + let session = sm + .create_session( + temp_dir.path().to_path_buf(), + "Same second truncation".to_string(), + SessionType::User, + GooseMode::default(), + ) + .await + .unwrap(); + + let timestamp = "2026-06-23T12:00:00Z"; + sm.add_message( + &session.id, + &Message::assistant() + .with_text("assistant reply") + .with_id("assistant"), + ) + .await + .unwrap(); + set_message_timestamp(&sm, &session.id, "assistant", timestamp).await; + + sm.add_message( + &session.id, + &Message::user() + .with_text("terminal history") + .with_id("terminal-history"), + ) + .await + .unwrap(); + set_message_timestamp(&sm, &session.id, "terminal-history", timestamp).await; + + sm.add_message( + &session.id, + &Message::user() + .with_text("next prompt") + .with_id("next-prompt"), + ) + .await + .unwrap(); + set_message_timestamp(&sm, &session.id, "next-prompt", timestamp).await; + + sm.truncate_conversation_from_message(&session.id, "terminal-history") + .await + .unwrap(); + + let reloaded = sm.get_session(&session.id, true).await.unwrap(); + let messages = reloaded.conversation.unwrap().messages().to_vec(); + assert_eq!(messages.len(), 1); + assert_eq!(messages[0].id.as_deref(), Some("assistant")); + assert_eq!(messages[0].as_concat_text(), "assistant reply"); + } + #[tokio::test] async fn test_maybe_update_name_updates_eligible_session() { let temp_dir = TempDir::new().unwrap(); diff --git a/documentation/docs/guides/goose-cli-commands.md b/documentation/docs/guides/goose-cli-commands.md index 29d5347fb5fc..09e1d09f2c88 100644 --- a/documentation/docs/guides/goose-cli-commands.md +++ b/documentation/docs/guides/goose-cli-commands.md @@ -351,13 +351,13 @@ goose session export --path ./my-session.jsonl -o exported.md --- #### session diagnostics [options] -Generate a comprehensive diagnostics bundle for troubleshooting issues with a specific session. +Generate a comprehensive diagnostics JSON report for troubleshooting issues with a specific session. **Options:** - **`--session-id `**: Generate diagnostics for a specific session by ID - **`-n, --name `**: Generate diagnostics for a specific session by name - **`--path `**: Generate diagnostics for a specific session by file path (legacy) -- **`-o, --output `**: Save diagnostics bundle to a specific file path (default: `diagnostics_{session_id}.zip`) +- **`-o, --output `**: Save diagnostics report to a specific file path (default: `diagnostics_{session_id}.json`) **What's included:** - **System Information**: App version, operating system, architecture, and timestamp @@ -374,18 +374,18 @@ goose session diagnostics --session-id 20251108_5 goose session diagnostics -n my-project-session # Save diagnostics to a custom location -goose session diagnostics --session-id 20251108_5 -o /path/to/my-diagnostics.zip +goose session diagnostics --session-id 20251108_5 -o /path/to/my-diagnostics.json # Interactive selection (prompts you to choose a session) goose session diagnostics ``` :::warning Privacy Notice -Diagnostics bundles contain your session messages and system information. If your session includes sensitive data (API keys, personal information, proprietary code), review the contents before sharing publicly. +Diagnostics reports contain your session messages and system information. If your session includes sensitive data (API keys, personal information, proprietary code), review the contents before sharing publicly. ::: :::tip -Generate diagnostics before reporting bugs to provide technical details that help with faster resolution. The ZIP file can be attached to GitHub issues or shared with support. +Generate diagnostics before reporting bugs to provide technical details that help with faster resolution. The JSON file can be attached to GitHub issues or shared with support. ::: --- diff --git a/documentation/docs/troubleshooting/diagnostics-and-reporting.md b/documentation/docs/troubleshooting/diagnostics-and-reporting.md index 084ee38749ed..20e69a82ccf0 100644 --- a/documentation/docs/troubleshooting/diagnostics-and-reporting.md +++ b/documentation/docs/troubleshooting/diagnostics-and-reporting.md @@ -12,13 +12,13 @@ goose provides several built-in features to help you get support, report issues, | Feature | Purpose | Location | Output | |---------|---------|----------|---------| -| **Diagnostics** | Generate troubleshooting data | Chat input toolbar | ZIP file with system info, logs, and session data | +| **Diagnostics** | Generate troubleshooting data | Chat input toolbar | JSON report with system info, logs, and session data | | **Report a Bug** | Submit bug reports | Chat input toolbar OR Settings → App → Help & feedback | Opens GitHub issue template | | **Request a Feature** | Suggest new features | Settings → App → Help & feedback | Opens GitHub issue template | ## Diagnostics System -The diagnostics feature creates a comprehensive troubleshooting bundle that includes system information, session data, configuration files, and recent logs. This is invaluable for debugging issues or getting technical support. +The diagnostics feature creates a comprehensive troubleshooting JSON report that includes system information, session data, configuration files, and recent logs. This is invaluable for debugging issues or getting technical support. ### Generating Diagnostics @@ -27,8 +27,10 @@ The diagnostics feature creates a comprehensive troubleshooting bundle that incl 1. In an active chat session, look for the icon in the bottom toolbar 2. Click the diagnostics button 3. Review the information in the modal about what data will be collected - 4. Click `Download` to generate and save the diagnostics bundle - 5. The ZIP file will be saved as `diagnostics_{session_id}.zip` + 4. Click `Download` to generate and save the diagnostics report + 5. The JSON file will be saved as `diagnostics_{session_id}.json` + + You can use `scripts/diagnostics-viewer.py` to inspect downloaded diagnostics reports; by default it looks in `~/Downloads`. :::tip The diagnostics button is only available when you have an active session, as it needs a session ID to generate the bundle. @@ -45,7 +47,7 @@ The diagnostics feature creates a comprehensive troubleshooting bundle that incl goose session diagnostics # Save to a custom location - goose session diagnostics --session-id --output /path/to/diagnostics.zip + goose session diagnostics --session-id --output /path/to/diagnostics.json ``` To find your session ID, first list available sessions: @@ -65,17 +67,18 @@ The diagnostics feature creates a comprehensive troubleshooting bundle that incl ### Using Diagnostics Data -The diagnostics ZIP file contains several folders: - -``` -diagnostics_abc123def.zip -├── logs/ -│ ├── goose-2024-01-15.jsonl -│ ├── goose-2024-01-14.jsonl -│ └── ... -├── session.json # Your session messages -├── config.yaml # Configuration files (if they exist) -└── system.txt # System information +The diagnostics JSON file contains structured sections: + +```json +{ + "system": {}, + "session": {}, + "config": {}, + "logs": {}, + "prompts": [], + "schedule": {}, + "errors": [] +} ``` **When to generate diagnostics:** @@ -154,4 +157,3 @@ For issues not resolved by diagnostics: - **[Session and System Logs](/docs/guides/logs)**: View detailed logs for debugging individual sessions - **[Telemetry Export](/docs/guides/environment-variables#observability)**: Configure telemetry for performance analysis and production monitoring - diff --git a/scripts/diagnostics-viewer.py b/scripts/diagnostics-viewer.py index 3129eadf892e..8ee4f92f9611 100755 --- a/scripts/diagnostics-viewer.py +++ b/scripts/diagnostics-viewer.py @@ -6,9 +6,9 @@ WARNING: entirely vibe coded. use as a throwaway tool -Diagnostics Viewer - Browse and inspect Goose diagnostics bundles. +Diagnostics Viewer - Browse and inspect Goose diagnostics reports. -Scans for diagnostics zip files, displays their sessions, and provides +Scans for diagnostics JSON reports and legacy zip files, displays their sessions, and provides an interactive viewer for examining session data, logs, and other files. """ import json @@ -188,34 +188,114 @@ def compose(self) -> ComposeResult: class DiagnosticsSession: - """Represents a diagnostics bundle.""" + """Represents a diagnostics report or legacy diagnostics bundle.""" - def __init__(self, zip_path: Path): - self.zip_path = zip_path + def __init__(self, path: Path): + self.path = path + self.is_zip = path.suffix == ".zip" self.name = "Unknown Session" - self.session_id = zip_path.stem - self.created_at = zip_path.stat().st_mtime + self.session_id = path.stem + self.created_at = path.stat().st_mtime + self.report = None self._load_session_name() def _load_session_name(self): - """Extract session name from session.json.""" + """Extract session name from the report.""" + if not self.is_zip: + self._load_json_report() + session = (self.report or {}).get("session") or {} + self.name = session.get("name", "Unknown Session") + self.session_id = session.get("id", self.path.stem) + return + try: - with zipfile.ZipFile(self.zip_path, 'r') as zf: + with zipfile.ZipFile(self.path, 'r') as zf: # Find session.json for name in zf.namelist(): if name.endswith('session.json'): with zf.open(name) as f: data = json.load(f) self.name = data.get('name', 'Unknown Session') - self.session_id = data.get('id', self.zip_path.stem) + self.session_id = data.get('id', self.path.stem) break except Exception as e: self.name = f"Error loading: {e}" + def _load_json_report(self): + if self.report is not None: + return + + try: + self.report = json.loads(self.path.read_text()) + except Exception as e: + self.report = {"error": f"Error loading: {e}"} + + def _json_virtual_files(self) -> dict[str, str]: + self._load_json_report() + report = self.report or {} + files = { + "diagnostics.json": json.dumps(report, indent=2), + } + + for key, filename in [ + ("system", "system.json"), + ("config", "config.json"), + ("extensions", "extensions.json"), + ("session", "session.json"), + ("schedule", "schedule.json"), + ("errors", "errors.json"), + ]: + value = report.get(key) + if value is not None: + files[filename] = json.dumps(value, indent=2) + + logs = report.get("logs") or {} + server = logs.get("server") + if isinstance(server, dict) and server.get("content") is not None: + files["logs/server.txt"] = server["content"] + + llm_logs = logs.get("llm") or [] + for index, entry in enumerate(llm_logs): + if isinstance(entry, dict) and entry.get("content") is not None: + path = Path(entry.get("path") or f"llm_request.{index}.jsonl") + files[f"logs/{path.name}"] = entry["content"] + + config = report.get("config") or {} + if isinstance(config, dict) and config.get("configYaml"): + files["config.yaml"] = config["configYaml"] + + for prompt in report.get("prompts") or []: + if isinstance(prompt, dict) and prompt.get("name") and prompt.get("content") is not None: + files[f"prompts/{prompt['name']}.txt"] = prompt["content"] + + for recipe in report.get("scheduledRecipes") or []: + if isinstance(recipe, dict) and recipe.get("path") and recipe.get("content") is not None: + path = Path(recipe["path"]) + files[f"scheduled_recipes/{path.name}"] = recipe["content"] + + return files + def get_file_list(self) -> list[str]: - """Get list of files in the zip, sorted with system.txt first.""" + """Get list of report files, sorted with system first.""" + if not self.is_zip: + files = list(self._json_virtual_files().keys()) + + def sort_key(f): + if f == "system.json": + return (0, f) + elif f == "session.json": + return (1, f) + elif f == "config.yaml" or f == "config.json": + return (2, f) + elif f == "diagnostics.json": + return (3, f) + else: + return (4, f) + + return sorted(files, key=sort_key) + try: - with zipfile.ZipFile(self.zip_path, 'r') as zf: + with zipfile.ZipFile(self.path, 'r') as zf: files = zf.namelist() # Sort: system.txt first, then session.json, then alphabetically @@ -234,13 +314,16 @@ def sort_key(f): return [] def read_file(self, filename: str) -> Optional[str]: - """Read a file from the zip. + """Read a file from the report. Returns: File content as string, or None if file cannot be read. """ + if not self.is_zip: + return self._json_virtual_files().get(filename) + try: - with zipfile.ZipFile(self.zip_path, 'r') as zf: + with zipfile.ZipFile(self.path, 'r') as zf: with zf.open(filename) as f: return f.read().decode('utf-8', errors='replace') except Exception: @@ -591,8 +674,8 @@ def on_mount(self): list_view = self.query_one(ListView) for session in self.sessions: item = ListItem( - Label(f"{session.name}\n[dim]{session.zip_path.name}[/dim]"), - name=session.zip_path.name + Label(f"{session.name}\n[dim]{session.path.name}[/dim]"), + name=session.path.name ) list_view.append(item) @@ -752,13 +835,14 @@ def on_mount(self): self.show_session_list() def scan_diagnostics(self): - """Scan for diagnostics zip files.""" + """Scan for diagnostics JSON reports and legacy zip files.""" self.sessions = [] - # Find all diagnostics zip files - for zip_path in self.diagnostics_dir.glob("diagnostics*.zip"): - session = DiagnosticsSession(zip_path) - self.sessions.append(session) + for path in [ + *self.diagnostics_dir.glob("diagnostics*.json"), + *self.diagnostics_dir.glob("diagnostics*.zip"), + ]: + self.sessions.append(DiagnosticsSession(path)) # Sort by creation time (newest first) self.sessions.sort(key=lambda s: s.created_at, reverse=True) @@ -781,9 +865,9 @@ def show_session_viewer(self, session: DiagnosticsSession): def on_list_view_selected(self, event: ListView.Selected): """Handle session selection.""" - # Find the session by zip name + # Find the session by diagnostics file name session_name = event.item.name - session = next((s for s in self.sessions if s.zip_path.name == session_name), None) + session = next((s for s in self.sessions if s.path.name == session_name), None) if session: self.show_session_viewer(session) diff --git a/ui/desktop/openapi.json b/ui/desktop/openapi.json index 1e7a991e423d..e299f4b33943 100644 --- a/ui/desktop/openapi.json +++ b/ui/desktop/openapi.json @@ -1730,6 +1730,19 @@ ], "operationId": "diagnostics", "parameters": [ + { + "name": "level", + "in": "query", + "required": false, + "schema": { + "allOf": [ + { + "$ref": "#/components/schemas/DiagnosticsLevel" + } + ], + "nullable": true + } + }, { "name": "session_id", "in": "path", @@ -1741,12 +1754,11 @@ ], "responses": { "200": { - "description": "Diagnostics zip file", + "description": "Diagnostics report", "content": { - "application/zip": { + "application/json": { "schema": { - "type": "string", - "format": "binary" + "$ref": "#/components/schemas/DiagnosticsReport" } } } @@ -4617,6 +4629,200 @@ } } }, + "DiagnosticsConfig": { + "type": "object", + "required": [ + "configPath", + "truncated" + ], + "properties": { + "configPath": { + "type": "string" + }, + "configYaml": { + "type": "string", + "nullable": true + }, + "truncated": { + "type": "boolean" + } + } + }, + "DiagnosticsError": { + "type": "object", + "required": [ + "message" + ], + "properties": { + "message": { + "type": "string" + }, + "path": { + "type": "string", + "nullable": true + } + } + }, + "DiagnosticsExtensions": { + "type": "object", + "required": [ + "enabled" + ], + "properties": { + "enabled": { + "type": "array", + "items": { + "type": "string" + } + } + } + }, + "DiagnosticsLevel": { + "type": "string", + "enum": [ + "summary", + "full" + ] + }, + "DiagnosticsLogs": { + "type": "object", + "required": [ + "llm" + ], + "properties": { + "llm": { + "type": "array", + "items": { + "$ref": "#/components/schemas/DiagnosticsTextFile" + } + }, + "server": { + "allOf": [ + { + "$ref": "#/components/schemas/DiagnosticsTextFile" + } + ], + "nullable": true + } + } + }, + "DiagnosticsPrompt": { + "type": "object", + "required": [ + "name", + "content" + ], + "properties": { + "content": { + "type": "string" + }, + "name": { + "type": "string" + } + } + }, + "DiagnosticsReport": { + "type": "object", + "required": [ + "schemaVersion", + "generatedAt", + "level", + "system", + "extensions", + "logs", + "prompts", + "scheduledRecipes", + "errors" + ], + "properties": { + "config": { + "allOf": [ + { + "$ref": "#/components/schemas/DiagnosticsConfig" + } + ], + "nullable": true + }, + "errors": { + "type": "array", + "items": { + "$ref": "#/components/schemas/DiagnosticsError" + } + }, + "extensions": { + "$ref": "#/components/schemas/DiagnosticsExtensions" + }, + "generatedAt": { + "type": "string" + }, + "level": { + "$ref": "#/components/schemas/DiagnosticsLevel" + }, + "logs": { + "$ref": "#/components/schemas/DiagnosticsLogs" + }, + "prompts": { + "type": "array", + "items": { + "$ref": "#/components/schemas/DiagnosticsPrompt" + } + }, + "schedule": { + "nullable": true + }, + "scheduledRecipes": { + "type": "array", + "items": { + "$ref": "#/components/schemas/DiagnosticsScheduledRecipe" + } + }, + "schemaVersion": { + "type": "integer", + "format": "int32", + "minimum": 0 + }, + "session": { + "nullable": true + }, + "system": { + "$ref": "#/components/schemas/SystemInfo" + } + } + }, + "DiagnosticsScheduledRecipe": { + "type": "object", + "required": [ + "path", + "content" + ], + "properties": { + "content": { + "type": "string" + }, + "path": { + "type": "string" + } + } + }, + "DiagnosticsTextFile": { + "type": "object", + "required": [ + "path", + "content", + "truncated" + ], + "properties": { + "content": { + "type": "string" + }, + "path": { + "type": "string" + }, + "truncated": { + "type": "boolean" + } + } + }, "DictationProvider": { "type": "string", "enum": [ diff --git a/ui/desktop/src/acp/diagnostics.ts b/ui/desktop/src/acp/diagnostics.ts new file mode 100644 index 000000000000..553ca8352c9a --- /dev/null +++ b/ui/desktop/src/acp/diagnostics.ts @@ -0,0 +1,16 @@ +import { getAcpClient } from './acpConnection'; +import type { DiagnosticsReport } from '../api'; + +export type DiagnosticsLevel = 'summary' | 'full'; + +export async function getDiagnosticsReport( + sessionId: string, + level: DiagnosticsLevel +): Promise { + const client = await getAcpClient(); + const response = await client.goose.diagnosticsGet_unstable({ + sessionId, + level, + }); + return response.report as DiagnosticsReport; +} diff --git a/ui/desktop/src/api/index.ts b/ui/desktop/src/api/index.ts index 17ed66891d55..33ddd72656c4 100644 --- a/ui/desktop/src/api/index.ts +++ b/ui/desktop/src/api/index.ts @@ -1,4 +1,4 @@ // This file is auto-generated by @hey-api/openapi-ts export { addExtension, agentAddExtension, agentRemoveExtension, callTool, cancelDownload, cancelLocalModelDownload, checkProvider, cleanupProviderCache, configureProviderOauth, confirmToolAction, createCustomProvider, createSchedule, decodeRecipe, deleteLocalModel, deleteModel, deleteProviderSecret, deleteRecipe, deleteSchedule, diagnostics, downloadHfModel, downloadModel, encodeRecipe, exportApp, forkSession, getCanonicalModelInfo, getCustomProvider, getDictationConfig, getDownloadProgress, getExtensions, getFeatures, getLocalModelDownloadProgress, getModelSettings, getPrompt, getPrompts, getProviderCatalog, getProviderCatalogTemplate, getProviderModelInfo, getProviderModels, getRepoFiles, getSession, getSessionExtensions, getSlashCommands, getTools, getTunnelStatus, importApp, importSessionNostr, inspectRunningJob, killRunningJob, listApps, listBuiltinChatTemplates, listLocalModels, listModels, listProviderSecrets, listRecipes, listSchedules, mcpUiProxy, type Options, parseRecipe, pauseSchedule, providers, readAllConfig, readConfig, readResource, recipeToYaml, removeConfig, removeCustomProvider, removeExtension, reply, resetPrompt, restartAgent, resumeAgent, runNowHandler, savePrompt, saveRecipe, scanRecipe, scheduleRecipe, searchHfModels, sendTelemetryEvent, sessionCancel, sessionEvents, sessionReply, sessionsHandler, setConfigProvider, setRecipeSlashCommand, shareSessionNostr, startAgent, startNanogptSetup, startOpenrouterSetup, startTetrateSetup, startTunnel, status, stopAgent, stopTunnel, syncFeaturedModels, systemInfo, transcribeDictation, unpauseSchedule, updateAgentProvider, updateCustomProvider, updateFromSession, updateModelSettings, updateSchedule, updateSession, updateSessionName, updateSessionUserRecipeValues, updateWorkingDir, upsertConfig, upsertPermissions, validateConfig } from './sdk.gen'; -export type { ActionRequired, ActionRequiredData, AddExtensionData, AddExtensionErrors, AddExtensionRequest, AddExtensionResponse, AddExtensionResponses, AgentAddExtensionData, AgentAddExtensionErrors, AgentAddExtensionResponse, AgentAddExtensionResponses, AgentRemoveExtensionData, AgentRemoveExtensionErrors, AgentRemoveExtensionResponse, AgentRemoveExtensionResponses, Annotations, Author, CallToolData, CallToolError, CallToolErrors, CallToolRequest, CallToolResponse, CallToolResponse2, CallToolResponses, CancelDownloadData, CancelDownloadErrors, CancelDownloadResponses, CancelLocalModelDownloadData, CancelLocalModelDownloadErrors, CancelLocalModelDownloadResponses, CancelRequest, ChatRequest, ChatTemplate, CheckProviderData, CheckProviderRequest, CleanupProviderCacheData, CleanupProviderCacheErrors, CleanupProviderCacheResponse, CleanupProviderCacheResponses, ClientOptions, CommandType, ConfigKey, ConfigKeyQuery, ConfigResponse, ConfigureProviderOauthData, ConfigureProviderOauthErrors, ConfigureProviderOauthResponses, ConfirmToolActionData, ConfirmToolActionErrors, ConfirmToolActionRequest, ConfirmToolActionResponses, Content, ContentBlock, Conversation, CreateCustomProviderData, CreateCustomProviderErrors, CreateCustomProviderResponse, CreateCustomProviderResponse2, CreateCustomProviderResponses, CreateScheduleData, CreateScheduleErrors, CreateScheduleRequest, CreateScheduleResponse, CreateScheduleResponses, CspMetadata, DeclarativeProviderConfig, DecodeRecipeData, DecodeRecipeErrors, DecodeRecipeRequest, DecodeRecipeResponse, DecodeRecipeResponse2, DecodeRecipeResponses, DeleteLocalModelData, DeleteLocalModelErrors, DeleteLocalModelResponses, DeleteModelData, DeleteModelErrors, DeleteModelResponses, DeleteProviderSecretData, DeleteProviderSecretErrors, DeleteProviderSecretResponse, DeleteProviderSecretResponses, DeleteRecipeData, DeleteRecipeErrors, DeleteRecipeRequest, DeleteRecipeResponse, DeleteRecipeResponses, DeleteScheduleData, DeleteScheduleErrors, DeleteScheduleResponse, DeleteScheduleResponses, DiagnosticsData, DiagnosticsErrors, DiagnosticsResponse, DiagnosticsResponses, DictationProvider, DictationProviderStatus, DownloadHfModelData, DownloadHfModelErrors, DownloadHfModelResponse, DownloadHfModelResponses, DownloadModelData, DownloadModelErrors, DownloadModelRequest, DownloadModelResponses, DownloadProgress, DownloadStatus, EmbeddedResource, EncodeRecipeData, EncodeRecipeErrors, EncodeRecipeRequest, EncodeRecipeResponse, EncodeRecipeResponse2, EncodeRecipeResponses, Envs, EnvVarConfig, ErrorResponse, ExportAppData, ExportAppError, ExportAppErrors, ExportAppResponse, ExportAppResponses, ExtensionConfig, ExtensionData, ExtensionEntry, ExtensionLoadResult, ExtensionQuery, ExtensionResponse, FeaturesResponse, ForkRequest, ForkResponse, ForkSessionData, ForkSessionErrors, ForkSessionResponse, ForkSessionResponses, FrontendToolRequest, GetCanonicalModelInfoData, GetCanonicalModelInfoResponse, GetCanonicalModelInfoResponses, GetCustomProviderData, GetCustomProviderErrors, GetCustomProviderResponse, GetCustomProviderResponses, GetDictationConfigData, GetDictationConfigResponse, GetDictationConfigResponses, GetDownloadProgressData, GetDownloadProgressErrors, GetDownloadProgressResponse, GetDownloadProgressResponses, GetExtensionsData, GetExtensionsErrors, GetExtensionsResponse, GetExtensionsResponses, GetFeaturesData, GetFeaturesResponse, GetFeaturesResponses, GetLocalModelDownloadProgressData, GetLocalModelDownloadProgressErrors, GetLocalModelDownloadProgressResponse, GetLocalModelDownloadProgressResponses, GetModelSettingsData, GetModelSettingsErrors, GetModelSettingsResponse, GetModelSettingsResponses, GetPromptData, GetPromptErrors, GetPromptResponse, GetPromptResponses, GetPromptsData, GetPromptsResponse, GetPromptsResponses, GetProviderCatalogData, GetProviderCatalogErrors, GetProviderCatalogResponse, GetProviderCatalogResponses, GetProviderCatalogTemplateData, GetProviderCatalogTemplateErrors, GetProviderCatalogTemplateResponse, GetProviderCatalogTemplateResponses, GetProviderModelInfoData, GetProviderModelInfoErrors, GetProviderModelInfoResponse, GetProviderModelInfoResponses, GetProviderModelsData, GetProviderModelsErrors, GetProviderModelsResponse, GetProviderModelsResponses, GetRepoFilesData, GetRepoFilesResponse, GetRepoFilesResponses, GetSessionData, GetSessionErrors, GetSessionExtensionsData, GetSessionExtensionsErrors, GetSessionExtensionsResponse, GetSessionExtensionsResponses, GetSessionResponse, GetSessionResponses, GetSlashCommandsData, GetSlashCommandsResponse, GetSlashCommandsResponses, GetToolsData, GetToolsErrors, GetToolsQuery, GetToolsResponse, GetToolsResponses, GetTunnelStatusData, GetTunnelStatusResponse, GetTunnelStatusResponses, GooseApp, GooseMode, HfGgufFile, HfModelInfo, HfModelVariant, HfQuantVariant, Icon, IconTheme, ImageContent, ImportAppData, ImportAppError, ImportAppErrors, ImportAppRequest, ImportAppResponse, ImportAppResponse2, ImportAppResponses, ImportSessionNostrData, ImportSessionNostrErrors, ImportSessionNostrRequest, ImportSessionNostrResponse, ImportSessionNostrResponses, InferenceMetadata, InspectJobResponse, InspectRunningJobData, InspectRunningJobErrors, InspectRunningJobResponse, InspectRunningJobResponses, JsonObject, KillJobResponse, KillRunningJobData, KillRunningJobResponses, ListAppsData, ListAppsError, ListAppsErrors, ListAppsRequest, ListAppsResponse, ListAppsResponse2, ListAppsResponses, ListBuiltinChatTemplatesData, ListBuiltinChatTemplatesResponse, ListBuiltinChatTemplatesResponses, ListLocalModelsData, ListLocalModelsResponse, ListLocalModelsResponses, ListModelsData, ListModelsResponse, ListModelsResponses, ListProviderSecretsData, ListProviderSecretsErrors, ListProviderSecretsResponse, ListProviderSecretsResponses, ListRecipeResponse, ListRecipesData, ListRecipesErrors, ListRecipesResponse, ListRecipesResponses, ListSchedulesData, ListSchedulesErrors, ListSchedulesResponse, ListSchedulesResponse2, ListSchedulesResponses, LoadedProvider, LocalModelResponse, McpAppResource, McpUiProxyData, McpUiProxyErrors, McpUiProxyResponses, Message, MessageContent, MessageEvent, MessageMetadata, ModelCapabilities, ModelConfig, ModelDownloadStatus, ModelInfo, ModelInfoData, ModelInfoQuery, ModelInfoResponse, ModelSettings, ModelTemplate, ParseRecipeData, ParseRecipeError, ParseRecipeErrors, ParseRecipeRequest, ParseRecipeResponse, ParseRecipeResponse2, ParseRecipeResponses, PauseScheduleData, PauseScheduleErrors, PauseScheduleResponse, PauseScheduleResponses, Permission, PermissionLevel, PermissionsMetadata, PrincipalType, PromptContentResponse, PromptsListResponse, ProviderCatalogEntry, ProviderDetails, ProviderEngine, ProviderMetadata, ProviderModelInfoQuery, ProvidersData, ProviderSecret, ProviderSecretsResponse, ProviderSecretStatus, ProviderSecretStorage, ProvidersResponse, ProvidersResponse2, ProvidersResponses, ProviderTemplate, ProviderType, RawAudioContent, RawEmbeddedResource, RawImageContent, RawResource, RawTextContent, ReadAllConfigData, ReadAllConfigResponse, ReadAllConfigResponses, ReadConfigData, ReadConfigErrors, ReadConfigResponses, ReadResourceData, ReadResourceErrors, ReadResourceRequest, ReadResourceResponse, ReadResourceResponse2, ReadResourceResponses, Recipe, RecipeManifest, RecipeParameter, RecipeParameterInputType, RecipeParameterRequirement, RecipeToYamlData, RecipeToYamlError, RecipeToYamlErrors, RecipeToYamlRequest, RecipeToYamlResponse, RecipeToYamlResponse2, RecipeToYamlResponses, RedactedThinkingContent, RemoveConfigData, RemoveConfigErrors, RemoveConfigResponse, RemoveConfigResponses, RemoveCustomProviderData, RemoveCustomProviderErrors, RemoveCustomProviderResponse, RemoveCustomProviderResponses, RemoveExtensionData, RemoveExtensionErrors, RemoveExtensionRequest, RemoveExtensionResponse, RemoveExtensionResponses, ReplyData, ReplyErrors, ReplyResponse, ReplyResponses, RepoVariantsResponse, ResetPromptData, ResetPromptErrors, ResetPromptResponse, ResetPromptResponses, ResourceContents, ResourceMetadata, Response, RestartAgentData, RestartAgentErrors, RestartAgentRequest, RestartAgentResponse, RestartAgentResponse2, RestartAgentResponses, ResumeAgentData, ResumeAgentErrors, ResumeAgentRequest, ResumeAgentResponse, ResumeAgentResponse2, ResumeAgentResponses, RetryConfig, Role, RunNowHandlerData, RunNowHandlerErrors, RunNowHandlerResponse, RunNowHandlerResponses, RunNowResponse, SamplingConfig, SavePromptData, SavePromptErrors, SavePromptRequest, SavePromptResponse, SavePromptResponses, SaveRecipeData, SaveRecipeError, SaveRecipeErrors, SaveRecipeRequest, SaveRecipeResponse, SaveRecipeResponse2, SaveRecipeResponses, ScanRecipeData, ScanRecipeRequest, ScanRecipeResponse, ScanRecipeResponse2, ScanRecipeResponses, ScheduledJob, ScheduleRecipeData, ScheduleRecipeErrors, ScheduleRecipeRequest, ScheduleRecipeResponses, SearchHfModelsData, SearchHfModelsErrors, SearchHfModelsResponse, SearchHfModelsResponses, SendTelemetryEventData, SendTelemetryEventResponses, Session, SessionCancelData, SessionCancelResponses, SessionDisplayInfo, SessionEventsData, SessionEventsErrors, SessionEventsResponse, SessionEventsResponses, SessionExtensionsResponse, SessionReplyData, SessionReplyErrors, SessionReplyRequest, SessionReplyResponse, SessionReplyResponse2, SessionReplyResponses, SessionsHandlerData, SessionsHandlerErrors, SessionsHandlerResponse, SessionsHandlerResponses, SessionsQuery, SessionType, SetConfigProviderData, SetProviderRequest, SetRecipeSlashCommandData, SetRecipeSlashCommandErrors, SetRecipeSlashCommandResponses, SetSlashCommandRequest, Settings, SetupResponse, ShareSessionNostrData, ShareSessionNostrErrors, ShareSessionNostrRequest, ShareSessionNostrResponse, ShareSessionNostrResponse2, ShareSessionNostrResponses, SlashCommand, SlashCommandsResponse, StartAgentData, StartAgentError, StartAgentErrors, StartAgentRequest, StartAgentResponse, StartAgentResponses, StartNanogptSetupData, StartNanogptSetupResponse, StartNanogptSetupResponses, StartOpenrouterSetupData, StartOpenrouterSetupResponse, StartOpenrouterSetupResponses, StartTetrateSetupData, StartTetrateSetupResponse, StartTetrateSetupResponses, StartTunnelData, StartTunnelError, StartTunnelErrors, StartTunnelResponse, StartTunnelResponses, StatusData, StatusResponse, StatusResponses, StopAgentData, StopAgentErrors, StopAgentRequest, StopAgentResponse, StopAgentResponses, StopTunnelData, StopTunnelError, StopTunnelErrors, StopTunnelResponses, SubRecipe, SuccessCheck, SyncFeaturedModelsData, SyncFeaturedModelsResponses, SystemInfo, SystemInfoData, SystemInfoResponse, SystemInfoResponses, SystemNotificationContent, SystemNotificationType, TaskSupport, TelemetryEventRequest, Template, TextContent, ThinkingContent, ThinkingEffort, TokenState, Tool, ToolAnnotations, ToolCallingMode, ToolConfirmationRequest, ToolExecution, ToolInfo, ToolPermission, ToolRequest, ToolResponse, TranscribeDictationData, TranscribeDictationErrors, TranscribeDictationResponse, TranscribeDictationResponses, TranscribeRequest, TranscribeResponse, TunnelInfo, TunnelState, UiMetadata, UnpauseScheduleData, UnpauseScheduleErrors, UnpauseScheduleResponse, UnpauseScheduleResponses, UpdateAgentProviderData, UpdateAgentProviderErrors, UpdateAgentProviderResponses, UpdateCustomProviderData, UpdateCustomProviderErrors, UpdateCustomProviderRequest, UpdateCustomProviderResponse, UpdateCustomProviderResponses, UpdateFromSessionData, UpdateFromSessionErrors, UpdateFromSessionRequest, UpdateFromSessionResponses, UpdateModelSettingsData, UpdateModelSettingsErrors, UpdateModelSettingsResponse, UpdateModelSettingsResponses, UpdateProviderRequest, UpdateScheduleData, UpdateScheduleErrors, UpdateScheduleRequest, UpdateScheduleResponse, UpdateScheduleResponses, UpdateSessionData, UpdateSessionErrors, UpdateSessionNameData, UpdateSessionNameErrors, UpdateSessionNameRequest, UpdateSessionNameResponses, UpdateSessionRequest, UpdateSessionResponses, UpdateSessionUserRecipeValuesData, UpdateSessionUserRecipeValuesError, UpdateSessionUserRecipeValuesErrors, UpdateSessionUserRecipeValuesRequest, UpdateSessionUserRecipeValuesResponse, UpdateSessionUserRecipeValuesResponse2, UpdateSessionUserRecipeValuesResponses, UpdateWorkingDirData, UpdateWorkingDirErrors, UpdateWorkingDirRequest, UpdateWorkingDirResponses, UpsertConfigData, UpsertConfigErrors, UpsertConfigQuery, UpsertConfigResponse, UpsertConfigResponses, UpsertPermissionsData, UpsertPermissionsErrors, UpsertPermissionsQuery, UpsertPermissionsResponse, UpsertPermissionsResponses, Usage, ValidateConfigData, ValidateConfigErrors, ValidateConfigResponse, ValidateConfigResponses, WhisperModelResponse, WindowProps } from './types.gen'; +export type { ActionRequired, ActionRequiredData, AddExtensionData, AddExtensionErrors, AddExtensionRequest, AddExtensionResponse, AddExtensionResponses, AgentAddExtensionData, AgentAddExtensionErrors, AgentAddExtensionResponse, AgentAddExtensionResponses, AgentRemoveExtensionData, AgentRemoveExtensionErrors, AgentRemoveExtensionResponse, AgentRemoveExtensionResponses, Annotations, Author, CallToolData, CallToolError, CallToolErrors, CallToolRequest, CallToolResponse, CallToolResponse2, CallToolResponses, CancelDownloadData, CancelDownloadErrors, CancelDownloadResponses, CancelLocalModelDownloadData, CancelLocalModelDownloadErrors, CancelLocalModelDownloadResponses, CancelRequest, ChatRequest, ChatTemplate, CheckProviderData, CheckProviderRequest, CleanupProviderCacheData, CleanupProviderCacheErrors, CleanupProviderCacheResponse, CleanupProviderCacheResponses, ClientOptions, CommandType, ConfigKey, ConfigKeyQuery, ConfigResponse, ConfigureProviderOauthData, ConfigureProviderOauthErrors, ConfigureProviderOauthResponses, ConfirmToolActionData, ConfirmToolActionErrors, ConfirmToolActionRequest, ConfirmToolActionResponses, Content, ContentBlock, Conversation, CreateCustomProviderData, CreateCustomProviderErrors, CreateCustomProviderResponse, CreateCustomProviderResponse2, CreateCustomProviderResponses, CreateScheduleData, CreateScheduleErrors, CreateScheduleRequest, CreateScheduleResponse, CreateScheduleResponses, CspMetadata, DeclarativeProviderConfig, DecodeRecipeData, DecodeRecipeErrors, DecodeRecipeRequest, DecodeRecipeResponse, DecodeRecipeResponse2, DecodeRecipeResponses, DeleteLocalModelData, DeleteLocalModelErrors, DeleteLocalModelResponses, DeleteModelData, DeleteModelErrors, DeleteModelResponses, DeleteProviderSecretData, DeleteProviderSecretErrors, DeleteProviderSecretResponse, DeleteProviderSecretResponses, DeleteRecipeData, DeleteRecipeErrors, DeleteRecipeRequest, DeleteRecipeResponse, DeleteRecipeResponses, DeleteScheduleData, DeleteScheduleErrors, DeleteScheduleResponse, DeleteScheduleResponses, DiagnosticsConfig, DiagnosticsData, DiagnosticsError, DiagnosticsErrors, DiagnosticsExtensions, DiagnosticsLevel, DiagnosticsLogs, DiagnosticsPrompt, DiagnosticsReport, DiagnosticsResponse, DiagnosticsResponses, DiagnosticsScheduledRecipe, DiagnosticsTextFile, DictationProvider, DictationProviderStatus, DownloadHfModelData, DownloadHfModelErrors, DownloadHfModelResponse, DownloadHfModelResponses, DownloadModelData, DownloadModelErrors, DownloadModelRequest, DownloadModelResponses, DownloadProgress, DownloadStatus, EmbeddedResource, EncodeRecipeData, EncodeRecipeErrors, EncodeRecipeRequest, EncodeRecipeResponse, EncodeRecipeResponse2, EncodeRecipeResponses, Envs, EnvVarConfig, ErrorResponse, ExportAppData, ExportAppError, ExportAppErrors, ExportAppResponse, ExportAppResponses, ExtensionConfig, ExtensionData, ExtensionEntry, ExtensionLoadResult, ExtensionQuery, ExtensionResponse, FeaturesResponse, ForkRequest, ForkResponse, ForkSessionData, ForkSessionErrors, ForkSessionResponse, ForkSessionResponses, FrontendToolRequest, GetCanonicalModelInfoData, GetCanonicalModelInfoResponse, GetCanonicalModelInfoResponses, GetCustomProviderData, GetCustomProviderErrors, GetCustomProviderResponse, GetCustomProviderResponses, GetDictationConfigData, GetDictationConfigResponse, GetDictationConfigResponses, GetDownloadProgressData, GetDownloadProgressErrors, GetDownloadProgressResponse, GetDownloadProgressResponses, GetExtensionsData, GetExtensionsErrors, GetExtensionsResponse, GetExtensionsResponses, GetFeaturesData, GetFeaturesResponse, GetFeaturesResponses, GetLocalModelDownloadProgressData, GetLocalModelDownloadProgressErrors, GetLocalModelDownloadProgressResponse, GetLocalModelDownloadProgressResponses, GetModelSettingsData, GetModelSettingsErrors, GetModelSettingsResponse, GetModelSettingsResponses, GetPromptData, GetPromptErrors, GetPromptResponse, GetPromptResponses, GetPromptsData, GetPromptsResponse, GetPromptsResponses, GetProviderCatalogData, GetProviderCatalogErrors, GetProviderCatalogResponse, GetProviderCatalogResponses, GetProviderCatalogTemplateData, GetProviderCatalogTemplateErrors, GetProviderCatalogTemplateResponse, GetProviderCatalogTemplateResponses, GetProviderModelInfoData, GetProviderModelInfoErrors, GetProviderModelInfoResponse, GetProviderModelInfoResponses, GetProviderModelsData, GetProviderModelsErrors, GetProviderModelsResponse, GetProviderModelsResponses, GetRepoFilesData, GetRepoFilesResponse, GetRepoFilesResponses, GetSessionData, GetSessionErrors, GetSessionExtensionsData, GetSessionExtensionsErrors, GetSessionExtensionsResponse, GetSessionExtensionsResponses, GetSessionResponse, GetSessionResponses, GetSlashCommandsData, GetSlashCommandsResponse, GetSlashCommandsResponses, GetToolsData, GetToolsErrors, GetToolsQuery, GetToolsResponse, GetToolsResponses, GetTunnelStatusData, GetTunnelStatusResponse, GetTunnelStatusResponses, GooseApp, GooseMode, HfGgufFile, HfModelInfo, HfModelVariant, HfQuantVariant, Icon, IconTheme, ImageContent, ImportAppData, ImportAppError, ImportAppErrors, ImportAppRequest, ImportAppResponse, ImportAppResponse2, ImportAppResponses, ImportSessionNostrData, ImportSessionNostrErrors, ImportSessionNostrRequest, ImportSessionNostrResponse, ImportSessionNostrResponses, InferenceMetadata, InspectJobResponse, InspectRunningJobData, InspectRunningJobErrors, InspectRunningJobResponse, InspectRunningJobResponses, JsonObject, KillJobResponse, KillRunningJobData, KillRunningJobResponses, ListAppsData, ListAppsError, ListAppsErrors, ListAppsRequest, ListAppsResponse, ListAppsResponse2, ListAppsResponses, ListBuiltinChatTemplatesData, ListBuiltinChatTemplatesResponse, ListBuiltinChatTemplatesResponses, ListLocalModelsData, ListLocalModelsResponse, ListLocalModelsResponses, ListModelsData, ListModelsResponse, ListModelsResponses, ListProviderSecretsData, ListProviderSecretsErrors, ListProviderSecretsResponse, ListProviderSecretsResponses, ListRecipeResponse, ListRecipesData, ListRecipesErrors, ListRecipesResponse, ListRecipesResponses, ListSchedulesData, ListSchedulesErrors, ListSchedulesResponse, ListSchedulesResponse2, ListSchedulesResponses, LoadedProvider, LocalModelResponse, McpAppResource, McpUiProxyData, McpUiProxyErrors, McpUiProxyResponses, Message, MessageContent, MessageEvent, MessageMetadata, ModelCapabilities, ModelConfig, ModelDownloadStatus, ModelInfo, ModelInfoData, ModelInfoQuery, ModelInfoResponse, ModelSettings, ModelTemplate, ParseRecipeData, ParseRecipeError, ParseRecipeErrors, ParseRecipeRequest, ParseRecipeResponse, ParseRecipeResponse2, ParseRecipeResponses, PauseScheduleData, PauseScheduleErrors, PauseScheduleResponse, PauseScheduleResponses, Permission, PermissionLevel, PermissionsMetadata, PrincipalType, PromptContentResponse, PromptsListResponse, ProviderCatalogEntry, ProviderDetails, ProviderEngine, ProviderMetadata, ProviderModelInfoQuery, ProvidersData, ProviderSecret, ProviderSecretsResponse, ProviderSecretStatus, ProviderSecretStorage, ProvidersResponse, ProvidersResponse2, ProvidersResponses, ProviderTemplate, ProviderType, RawAudioContent, RawEmbeddedResource, RawImageContent, RawResource, RawTextContent, ReadAllConfigData, ReadAllConfigResponse, ReadAllConfigResponses, ReadConfigData, ReadConfigErrors, ReadConfigResponses, ReadResourceData, ReadResourceErrors, ReadResourceRequest, ReadResourceResponse, ReadResourceResponse2, ReadResourceResponses, Recipe, RecipeManifest, RecipeParameter, RecipeParameterInputType, RecipeParameterRequirement, RecipeToYamlData, RecipeToYamlError, RecipeToYamlErrors, RecipeToYamlRequest, RecipeToYamlResponse, RecipeToYamlResponse2, RecipeToYamlResponses, RedactedThinkingContent, RemoveConfigData, RemoveConfigErrors, RemoveConfigResponse, RemoveConfigResponses, RemoveCustomProviderData, RemoveCustomProviderErrors, RemoveCustomProviderResponse, RemoveCustomProviderResponses, RemoveExtensionData, RemoveExtensionErrors, RemoveExtensionRequest, RemoveExtensionResponse, RemoveExtensionResponses, ReplyData, ReplyErrors, ReplyResponse, ReplyResponses, RepoVariantsResponse, ResetPromptData, ResetPromptErrors, ResetPromptResponse, ResetPromptResponses, ResourceContents, ResourceMetadata, Response, RestartAgentData, RestartAgentErrors, RestartAgentRequest, RestartAgentResponse, RestartAgentResponse2, RestartAgentResponses, ResumeAgentData, ResumeAgentErrors, ResumeAgentRequest, ResumeAgentResponse, ResumeAgentResponse2, ResumeAgentResponses, RetryConfig, Role, RunNowHandlerData, RunNowHandlerErrors, RunNowHandlerResponse, RunNowHandlerResponses, RunNowResponse, SamplingConfig, SavePromptData, SavePromptErrors, SavePromptRequest, SavePromptResponse, SavePromptResponses, SaveRecipeData, SaveRecipeError, SaveRecipeErrors, SaveRecipeRequest, SaveRecipeResponse, SaveRecipeResponse2, SaveRecipeResponses, ScanRecipeData, ScanRecipeRequest, ScanRecipeResponse, ScanRecipeResponse2, ScanRecipeResponses, ScheduledJob, ScheduleRecipeData, ScheduleRecipeErrors, ScheduleRecipeRequest, ScheduleRecipeResponses, SearchHfModelsData, SearchHfModelsErrors, SearchHfModelsResponse, SearchHfModelsResponses, SendTelemetryEventData, SendTelemetryEventResponses, Session, SessionCancelData, SessionCancelResponses, SessionDisplayInfo, SessionEventsData, SessionEventsErrors, SessionEventsResponse, SessionEventsResponses, SessionExtensionsResponse, SessionReplyData, SessionReplyErrors, SessionReplyRequest, SessionReplyResponse, SessionReplyResponse2, SessionReplyResponses, SessionsHandlerData, SessionsHandlerErrors, SessionsHandlerResponse, SessionsHandlerResponses, SessionsQuery, SessionType, SetConfigProviderData, SetProviderRequest, SetRecipeSlashCommandData, SetRecipeSlashCommandErrors, SetRecipeSlashCommandResponses, SetSlashCommandRequest, Settings, SetupResponse, ShareSessionNostrData, ShareSessionNostrErrors, ShareSessionNostrRequest, ShareSessionNostrResponse, ShareSessionNostrResponse2, ShareSessionNostrResponses, SlashCommand, SlashCommandsResponse, StartAgentData, StartAgentError, StartAgentErrors, StartAgentRequest, StartAgentResponse, StartAgentResponses, StartNanogptSetupData, StartNanogptSetupResponse, StartNanogptSetupResponses, StartOpenrouterSetupData, StartOpenrouterSetupResponse, StartOpenrouterSetupResponses, StartTetrateSetupData, StartTetrateSetupResponse, StartTetrateSetupResponses, StartTunnelData, StartTunnelError, StartTunnelErrors, StartTunnelResponse, StartTunnelResponses, StatusData, StatusResponse, StatusResponses, StopAgentData, StopAgentErrors, StopAgentRequest, StopAgentResponse, StopAgentResponses, StopTunnelData, StopTunnelError, StopTunnelErrors, StopTunnelResponses, SubRecipe, SuccessCheck, SyncFeaturedModelsData, SyncFeaturedModelsResponses, SystemInfo, SystemInfoData, SystemInfoResponse, SystemInfoResponses, SystemNotificationContent, SystemNotificationType, TaskSupport, TelemetryEventRequest, Template, TextContent, ThinkingContent, ThinkingEffort, TokenState, Tool, ToolAnnotations, ToolCallingMode, ToolConfirmationRequest, ToolExecution, ToolInfo, ToolPermission, ToolRequest, ToolResponse, TranscribeDictationData, TranscribeDictationErrors, TranscribeDictationResponse, TranscribeDictationResponses, TranscribeRequest, TranscribeResponse, TunnelInfo, TunnelState, UiMetadata, UnpauseScheduleData, UnpauseScheduleErrors, UnpauseScheduleResponse, UnpauseScheduleResponses, UpdateAgentProviderData, UpdateAgentProviderErrors, UpdateAgentProviderResponses, UpdateCustomProviderData, UpdateCustomProviderErrors, UpdateCustomProviderRequest, UpdateCustomProviderResponse, UpdateCustomProviderResponses, UpdateFromSessionData, UpdateFromSessionErrors, UpdateFromSessionRequest, UpdateFromSessionResponses, UpdateModelSettingsData, UpdateModelSettingsErrors, UpdateModelSettingsResponse, UpdateModelSettingsResponses, UpdateProviderRequest, UpdateScheduleData, UpdateScheduleErrors, UpdateScheduleRequest, UpdateScheduleResponse, UpdateScheduleResponses, UpdateSessionData, UpdateSessionErrors, UpdateSessionNameData, UpdateSessionNameErrors, UpdateSessionNameRequest, UpdateSessionNameResponses, UpdateSessionRequest, UpdateSessionResponses, UpdateSessionUserRecipeValuesData, UpdateSessionUserRecipeValuesError, UpdateSessionUserRecipeValuesErrors, UpdateSessionUserRecipeValuesRequest, UpdateSessionUserRecipeValuesResponse, UpdateSessionUserRecipeValuesResponse2, UpdateSessionUserRecipeValuesResponses, UpdateWorkingDirData, UpdateWorkingDirErrors, UpdateWorkingDirRequest, UpdateWorkingDirResponses, UpsertConfigData, UpsertConfigErrors, UpsertConfigQuery, UpsertConfigResponse, UpsertConfigResponses, UpsertPermissionsData, UpsertPermissionsErrors, UpsertPermissionsQuery, UpsertPermissionsResponse, UpsertPermissionsResponses, Usage, ValidateConfigData, ValidateConfigErrors, ValidateConfigResponse, ValidateConfigResponses, WhisperModelResponse, WindowProps } from './types.gen'; diff --git a/ui/desktop/src/api/types.gen.ts b/ui/desktop/src/api/types.gen.ts index b21738aed45b..7f8eb3e44fd8 100644 --- a/ui/desktop/src/api/types.gen.ts +++ b/ui/desktop/src/api/types.gen.ts @@ -247,6 +247,59 @@ export type DeleteRecipeRequest = { id: string; }; +export type DiagnosticsConfig = { + configPath: string; + configYaml?: string | null; + truncated: boolean; +}; + +export type DiagnosticsError = { + message: string; + path?: string | null; +}; + +export type DiagnosticsExtensions = { + enabled: Array; +}; + +export type DiagnosticsLevel = 'summary' | 'full'; + +export type DiagnosticsLogs = { + llm: Array; + server?: DiagnosticsTextFile | null; +}; + +export type DiagnosticsPrompt = { + content: string; + name: string; +}; + +export type DiagnosticsReport = { + config?: DiagnosticsConfig | null; + errors: Array; + extensions: DiagnosticsExtensions; + generatedAt: string; + level: DiagnosticsLevel; + logs: DiagnosticsLogs; + prompts: Array; + schedule?: unknown; + scheduledRecipes: Array; + schemaVersion: number; + session?: unknown; + system: SystemInfo; +}; + +export type DiagnosticsScheduledRecipe = { + content: string; + path: string; +}; + +export type DiagnosticsTextFile = { + content: string; + path: string; + truncated: boolean; +}; + export type DictationProvider = 'openai' | 'elevenlabs' | 'groq' | 'local'; export type DictationProviderStatus = { @@ -3098,7 +3151,9 @@ export type DiagnosticsData = { path: { session_id: string; }; - query?: never; + query?: { + level?: DiagnosticsLevel | null; + }; url: '/diagnostics/{session_id}'; }; @@ -3111,9 +3166,9 @@ export type DiagnosticsErrors = { export type DiagnosticsResponses = { /** - * Diagnostics zip file + * Diagnostics report */ - 200: Blob | File; + 200: DiagnosticsReport; }; export type DiagnosticsResponse = DiagnosticsResponses[keyof DiagnosticsResponses]; diff --git a/ui/desktop/src/components/ui/Diagnostics.tsx b/ui/desktop/src/components/ui/Diagnostics.tsx index 9ea9cc0f03ef..2951b1902d86 100644 --- a/ui/desktop/src/components/ui/Diagnostics.tsx +++ b/ui/desktop/src/components/ui/Diagnostics.tsx @@ -2,8 +2,8 @@ import React, { useState } from 'react'; import { AlertTriangle, Download, Github } from 'lucide-react'; import { Button } from './button'; import { toastError } from '../../toasts'; -import { diagnostics, systemInfo } from '../../api'; import { defineMessages, useIntl } from '../../i18n'; +import { getDiagnosticsReport } from '../../acp/diagnostics'; const i18n = defineMessages({ reportProblem: { @@ -13,7 +13,7 @@ const i18n = defineMessages({ description: { id: 'diagnosticsModal.description', defaultMessage: - 'You can download a diagnostics zip file to share with the team, or file a bug directly on GitHub with your system details pre-filled. A diagnostics report contains the following:', + 'You can download a diagnostics JSON report to share with the team, or file a bug directly on GitHub with your system details pre-filled. A diagnostics report contains the following:', }, systemInfo: { id: 'diagnosticsModal.systemInfo', @@ -66,7 +66,7 @@ const i18n = defineMessages({ }, diagnosticsErrorMsg: { id: 'diagnosticsModal.diagnosticsErrorMsg', - defaultMessage: 'Failed to download diagnostics', + defaultMessage: 'Failed to download diagnostics report', }, systemInfoErrorTitle: { id: 'diagnosticsModal.systemInfoErrorTitle', @@ -97,16 +97,14 @@ export const DiagnosticsModal: React.FC = ({ setIsDownloading(true); try { - const response = await diagnostics({ - path: { session_id: sessionId }, - throwOnError: true, + const report = await getDiagnosticsReport(sessionId, 'full'); + const blob = new Blob([`${JSON.stringify(report, null, 2)}\n`], { + type: 'application/json', }); - - const blob = new Blob([response.data], { type: 'application/zip' }); const url = window.URL.createObjectURL(blob); const a = document.createElement('a'); a.href = url; - a.download = `diagnostics_${sessionId}.zip`; + a.download = `diagnostics_${sessionId}.json`; document.body.appendChild(a); a.click(); document.body.removeChild(a); @@ -127,12 +125,12 @@ export const DiagnosticsModal: React.FC = ({ setIsFilingBug(true); try { - const response = await systemInfo({ throwOnError: true }); - const info = response.data; + const report = await getDiagnosticsReport(sessionId, 'summary'); + const info = report.system; const providerModel = info.provider && info.model - ? `${info.provider} – ${info.model}` + ? `${info.provider} - ${info.model}` : info.provider || info.model || '[e.g. Google – gemini-1.5-pro]'; const extensions = @@ -145,7 +143,7 @@ export const DiagnosticsModal: React.FC = ({ 💡 Before filing, please check common issues: https://goose-docs.ai/docs/troubleshooting -📦 To help us debug faster, attach your **diagnostics zip** if possible. +📦 To help us debug faster, attach your **diagnostics JSON report** if possible. 👉 How to capture it: https://goose-docs.ai/docs/troubleshooting/diagnostics-and-reporting/ A clear and concise description of what the bug is. diff --git a/ui/desktop/src/i18n/messages/en.json b/ui/desktop/src/i18n/messages/en.json index 8b19ba6382a9..122b6c7127ab 100644 --- a/ui/desktop/src/i18n/messages/en.json +++ b/ui/desktop/src/i18n/messages/en.json @@ -804,10 +804,10 @@ "defaultMessage": "Configuration settings" }, "diagnosticsModal.description": { - "defaultMessage": "You can download a diagnostics zip file to share with the team, or file a bug directly on GitHub with your system details pre-filled. A diagnostics report contains the following:" + "defaultMessage": "You can download a diagnostics JSON report to share with the team, or file a bug directly on GitHub with your system details pre-filled. A diagnostics report contains the following:" }, "diagnosticsModal.diagnosticsErrorMsg": { - "defaultMessage": "Failed to download diagnostics" + "defaultMessage": "Failed to download diagnostics report" }, "diagnosticsModal.diagnosticsErrorTitle": { "defaultMessage": "Diagnostics Error" diff --git a/ui/sdk/src/generated/client.gen.ts b/ui/sdk/src/generated/client.gen.ts index 7588717d4c51..fa13ca5e8abf 100644 --- a/ui/sdk/src/generated/client.gen.ts +++ b/ui/sdk/src/generated/client.gen.ts @@ -30,6 +30,8 @@ import type { DeleteRecipeRequest_unstable, DeleteSessionRequest, DeleteSourceRequest_unstable, + DiagnosticsGetRequest_unstable, + DiagnosticsGetResponse_unstable, DictationConfigRequest_unstable, DictationConfigResponse_unstable, DictationModelCancelRequest_unstable, @@ -135,6 +137,7 @@ import { zCustomProviderUpdateResponse_unstable, zDecodeRecipeResponse_unstable, zDefaultsReadResponse_unstable, + zDiagnosticsGetResponse_unstable, zDictationConfigResponse_unstable, zDictationModelDownloadProgressResponse_unstable, zDictationModelsListResponse_unstable, @@ -251,6 +254,18 @@ export class GooseExtClient { ) as SteerSessionResponse_unstable; } + async diagnosticsGet_unstable( + params: DiagnosticsGetRequest_unstable, + ): Promise { + const raw = await this.conn.extMethod( + "_goose/unstable/diagnostics/get", + params, + ); + return zDiagnosticsGetResponse_unstable.parse( + raw, + ) as DiagnosticsGetResponse_unstable; + } + async sessionDelete(params: DeleteSessionRequest): Promise { await this.conn.extMethod("session/delete", params); } diff --git a/ui/sdk/src/generated/index.ts b/ui/sdk/src/generated/index.ts index 543ba4e95d54..fc97f037657f 100644 --- a/ui/sdk/src/generated/index.ts +++ b/ui/sdk/src/generated/index.ts @@ -1,6 +1,6 @@ // This file is auto-generated by @hey-api/openapi-ts -export type { AddConfigExtensionRequest_unstable, AddSessionExtensionRequest_unstable, Annotations, ArchiveSessionRequest_unstable, AudioContent, BlobResourceContents, ContentBlock, CreateSourceRequest_unstable, CreateSourceResponse_unstable, CustomProviderConfigDto, CustomProviderCreateRequest_unstable, CustomProviderCreateResponse_unstable, CustomProviderDeleteRequest_unstable, CustomProviderDeleteResponse_unstable, CustomProviderReadRequest_unstable, CustomProviderReadResponse_unstable, CustomProviderUpdateRequest_unstable, CustomProviderUpdateResponse_unstable, DecodeRecipeRequest_unstable, DecodeRecipeResponse_unstable, DefaultsReadRequest_unstable, DefaultsReadResponse_unstable, DefaultsSaveRequest_unstable, DeleteRecipeRequest_unstable, DeleteSessionRequest, DeleteSourceRequest_unstable, DictationConfigRequest_unstable, DictationConfigResponse_unstable, DictationDownloadProgress, DictationLocalModelStatus, DictationModelCancelRequest_unstable, DictationModelDeleteRequest_unstable, DictationModelDownloadProgressRequest_unstable, DictationModelDownloadProgressResponse_unstable, DictationModelDownloadRequest_unstable, DictationModelOption, DictationModelSelectRequest_unstable, DictationModelsListRequest_unstable, DictationModelsListResponse_unstable, DictationProviderStatusEntry, DictationSecretDeleteRequest_unstable, DictationSecretSaveRequest_unstable, DictationTranscribeRequest_unstable, DictationTranscribeResponse_unstable, EmbeddedResource, EmbeddedResourceResource, EmptyResponse, EncodeRecipeRequest_unstable, EncodeRecipeResponse_unstable, EnvVariable, ExportSessionRequest_unstable, ExportSessionResponse_unstable, ExportSourceRequest_unstable, ExportSourceResponse_unstable, ExtAgentRequest, ExtAgentResponse, ExtNotification, ExtRequest, ExtResponse, GetAvailableExtensionsRequest_unstable, GetAvailableExtensionsResponse_unstable, GetConfigExtensionsRequest_unstable, GetConfigExtensionsResponse_unstable, GetSessionExtensionsRequest_unstable, GetSessionExtensionsResponse_unstable, GetSessionInfoRequest_unstable, GetSessionInfoResponse_unstable, GetToolsRequest_unstable, GetToolsResponse_unstable, GooseExtension, GooseExtensionEntry, GooseSessionNotification_unstable, GooseSessionUpdate, GooseToolCallRequest_unstable, GooseToolCallResponse_unstable, HttpHeader, ImageContent, ImportSessionRequest_unstable, ImportSessionResponse_unstable, ImportSourcesRequest_unstable, ImportSourcesResponse_unstable, ListProvidersRequest_unstable, ListProvidersResponse_unstable, ListRecipesRequest_unstable, ListRecipesResponse_unstable, ListSourcesRequest_unstable, ListSourcesResponse_unstable, McpServer, McpServerHttp, McpServerSse, McpServerStdio, OnboardingImportApplyRequest_unstable, OnboardingImportApplyResponse_unstable, OnboardingImportCandidate, OnboardingImportCounts, OnboardingImportScanRequest_unstable, OnboardingImportScanResponse_unstable, OnboardingImportSourceKind, ParseRecipeRequest_unstable, ParseRecipeResponse_unstable, PreferenceKey, PreferencesReadRequest_unstable, PreferencesReadResponse_unstable, PreferencesRemoveRequest_unstable, PreferencesSaveRequest_unstable, PreferenceValue, ProviderCatalogListRequest_unstable, ProviderCatalogListResponse_unstable, ProviderCatalogTemplateRequest_unstable, ProviderCatalogTemplateResponse_unstable, ProviderConfigAuthenticateRequest_unstable, ProviderConfigChangeResponse_unstable, ProviderConfigDeleteRequest_unstable, ProviderConfigFieldUpdate, ProviderConfigFieldValueDto, ProviderConfigKey, ProviderConfigReadRequest_unstable, ProviderConfigReadResponse_unstable, ProviderConfigSaveRequest_unstable, ProviderConfigStatusDto, ProviderConfigStatusRequest_unstable, ProviderConfigStatusResponse_unstable, ProviderInventoryEntryDto, ProviderInventoryModelDto, ProviderSetupCatalogEntryDto, ProviderSetupCatalogListRequest_unstable, ProviderSetupCatalogListResponse_unstable, ProviderSetupCategoryDto, ProviderSetupFieldDto, ProviderSetupGroupDto, ProviderSetupMethodDto, ProviderSupportedModelsListRequest_unstable, ProviderSupportedModelsListResponse_unstable, ProviderTemplateCapabilitiesDto, ProviderTemplateCatalogEntryDto, ProviderTemplateDto, ProviderTemplateModelDto, ReadResourceRequest_unstable, ReadResourceResponse_unstable, RecipeAuthorDto, RecipeDto, RecipeExtensionDto, RecipeListEntryDto, RecipeParameterDto, RecipeParameterInputTypeDto, RecipeParameterRequirementDto, RecipeParamsAction, RecipeParamsResponse_unstable, RecipeResponseDto, RecipeRetryConfigDto, RecipeSettingsDto, RecipeSuccessCheckDto, RecipeToYamlRequest_unstable, RecipeToYamlResponse_unstable, RefreshProviderInventoryRequest_unstable, RefreshProviderInventoryResponse_unstable, RefreshProviderInventorySkipDto, RefreshProviderInventorySkipReasonDto, RemoveConfigExtensionRequest_unstable, RemoveSessionExtensionRequest_unstable, RenameSessionRequest_unstable, RequestRecipeParams_unstable, ResourceLink, Role, SaveRecipeRequest_unstable, SaveRecipeResponse_unstable, ScanRecipeRequest_unstable, ScanRecipeResponse_unstable, ScheduleRecipeRequest_unstable, SessionId, SessionInfo, SessionSystemPromptMode, SessionUsageUpdate, SetConfigExtensionEnabledRequest_unstable, SetRecipeSlashCommandRequest_unstable, SetSessionSystemPromptRequest_unstable, SourceEntry, SourceScope, SourceType, StatusMessage, StatusMessageUpdate, SteerSessionRequest_unstable, SteerSessionResponse_unstable, SubRecipeDto, TextContent, TextResourceContents, TruncateSessionConversationRequest_unstable, UnarchiveSessionRequest_unstable, UpdateSessionProjectRequest_unstable, UpdateSourceRequest_unstable, UpdateSourceResponse_unstable, UpdateWorkingDirRequest_unstable } from './types.gen.js'; +export type { AddConfigExtensionRequest_unstable, AddSessionExtensionRequest_unstable, Annotations, ArchiveSessionRequest_unstable, AudioContent, BlobResourceContents, ContentBlock, CreateSourceRequest_unstable, CreateSourceResponse_unstable, CustomProviderConfigDto, CustomProviderCreateRequest_unstable, CustomProviderCreateResponse_unstable, CustomProviderDeleteRequest_unstable, CustomProviderDeleteResponse_unstable, CustomProviderReadRequest_unstable, CustomProviderReadResponse_unstable, CustomProviderUpdateRequest_unstable, CustomProviderUpdateResponse_unstable, DecodeRecipeRequest_unstable, DecodeRecipeResponse_unstable, DefaultsReadRequest_unstable, DefaultsReadResponse_unstable, DefaultsSaveRequest_unstable, DeleteRecipeRequest_unstable, DeleteSessionRequest, DeleteSourceRequest_unstable, DiagnosticsGetRequest_unstable, DiagnosticsGetResponse_unstable, DiagnosticsReportLevel, DictationConfigRequest_unstable, DictationConfigResponse_unstable, DictationDownloadProgress, DictationLocalModelStatus, DictationModelCancelRequest_unstable, DictationModelDeleteRequest_unstable, DictationModelDownloadProgressRequest_unstable, DictationModelDownloadProgressResponse_unstable, DictationModelDownloadRequest_unstable, DictationModelOption, DictationModelSelectRequest_unstable, DictationModelsListRequest_unstable, DictationModelsListResponse_unstable, DictationProviderStatusEntry, DictationSecretDeleteRequest_unstable, DictationSecretSaveRequest_unstable, DictationTranscribeRequest_unstable, DictationTranscribeResponse_unstable, EmbeddedResource, EmbeddedResourceResource, EmptyResponse, EncodeRecipeRequest_unstable, EncodeRecipeResponse_unstable, EnvVariable, ExportSessionRequest_unstable, ExportSessionResponse_unstable, ExportSourceRequest_unstable, ExportSourceResponse_unstable, ExtAgentRequest, ExtAgentResponse, ExtNotification, ExtRequest, ExtResponse, GetAvailableExtensionsRequest_unstable, GetAvailableExtensionsResponse_unstable, GetConfigExtensionsRequest_unstable, GetConfigExtensionsResponse_unstable, GetSessionExtensionsRequest_unstable, GetSessionExtensionsResponse_unstable, GetSessionInfoRequest_unstable, GetSessionInfoResponse_unstable, GetToolsRequest_unstable, GetToolsResponse_unstable, GooseExtension, GooseExtensionEntry, GooseSessionNotification_unstable, GooseSessionUpdate, GooseToolCallRequest_unstable, GooseToolCallResponse_unstable, HttpHeader, ImageContent, ImportSessionRequest_unstable, ImportSessionResponse_unstable, ImportSourcesRequest_unstable, ImportSourcesResponse_unstable, ListProvidersRequest_unstable, ListProvidersResponse_unstable, ListRecipesRequest_unstable, ListRecipesResponse_unstable, ListSourcesRequest_unstable, ListSourcesResponse_unstable, McpServer, McpServerHttp, McpServerSse, McpServerStdio, OnboardingImportApplyRequest_unstable, OnboardingImportApplyResponse_unstable, OnboardingImportCandidate, OnboardingImportCounts, OnboardingImportScanRequest_unstable, OnboardingImportScanResponse_unstable, OnboardingImportSourceKind, ParseRecipeRequest_unstable, ParseRecipeResponse_unstable, PreferenceKey, PreferencesReadRequest_unstable, PreferencesReadResponse_unstable, PreferencesRemoveRequest_unstable, PreferencesSaveRequest_unstable, PreferenceValue, ProviderCatalogListRequest_unstable, ProviderCatalogListResponse_unstable, ProviderCatalogTemplateRequest_unstable, ProviderCatalogTemplateResponse_unstable, ProviderConfigAuthenticateRequest_unstable, ProviderConfigChangeResponse_unstable, ProviderConfigDeleteRequest_unstable, ProviderConfigFieldUpdate, ProviderConfigFieldValueDto, ProviderConfigKey, ProviderConfigReadRequest_unstable, ProviderConfigReadResponse_unstable, ProviderConfigSaveRequest_unstable, ProviderConfigStatusDto, ProviderConfigStatusRequest_unstable, ProviderConfigStatusResponse_unstable, ProviderInventoryEntryDto, ProviderInventoryModelDto, ProviderSetupCatalogEntryDto, ProviderSetupCatalogListRequest_unstable, ProviderSetupCatalogListResponse_unstable, ProviderSetupCategoryDto, ProviderSetupFieldDto, ProviderSetupGroupDto, ProviderSetupMethodDto, ProviderSupportedModelsListRequest_unstable, ProviderSupportedModelsListResponse_unstable, ProviderTemplateCapabilitiesDto, ProviderTemplateCatalogEntryDto, ProviderTemplateDto, ProviderTemplateModelDto, ReadResourceRequest_unstable, ReadResourceResponse_unstable, RecipeAuthorDto, RecipeDto, RecipeExtensionDto, RecipeListEntryDto, RecipeParameterDto, RecipeParameterInputTypeDto, RecipeParameterRequirementDto, RecipeParamsAction, RecipeParamsResponse_unstable, RecipeResponseDto, RecipeRetryConfigDto, RecipeSettingsDto, RecipeSuccessCheckDto, RecipeToYamlRequest_unstable, RecipeToYamlResponse_unstable, RefreshProviderInventoryRequest_unstable, RefreshProviderInventoryResponse_unstable, RefreshProviderInventorySkipDto, RefreshProviderInventorySkipReasonDto, RemoveConfigExtensionRequest_unstable, RemoveSessionExtensionRequest_unstable, RenameSessionRequest_unstable, RequestRecipeParams_unstable, ResourceLink, Role, SaveRecipeRequest_unstable, SaveRecipeResponse_unstable, ScanRecipeRequest_unstable, ScanRecipeResponse_unstable, ScheduleRecipeRequest_unstable, SessionId, SessionInfo, SessionSystemPromptMode, SessionUsageUpdate, SetConfigExtensionEnabledRequest_unstable, SetRecipeSlashCommandRequest_unstable, SetSessionSystemPromptRequest_unstable, SourceEntry, SourceScope, SourceType, StatusMessage, StatusMessageUpdate, SteerSessionRequest_unstable, SteerSessionResponse_unstable, SubRecipeDto, TextContent, TextResourceContents, TruncateSessionConversationRequest_unstable, UnarchiveSessionRequest_unstable, UpdateSessionProjectRequest_unstable, UpdateSourceRequest_unstable, UpdateSourceResponse_unstable, UpdateWorkingDirRequest_unstable } from './types.gen.js'; export const GOOSE_EXT_METHODS = [ { @@ -43,6 +43,11 @@ export const GOOSE_EXT_METHODS = [ requestType: "SteerSessionRequest_unstable", responseType: "SteerSessionResponse_unstable", }, + { + method: "_goose/unstable/diagnostics/get", + requestType: "DiagnosticsGetRequest_unstable", + responseType: "DiagnosticsGetResponse_unstable", + }, { method: "session/delete", requestType: "DeleteSessionRequest", diff --git a/ui/sdk/src/generated/types.gen.ts b/ui/sdk/src/generated/types.gen.ts index 555289163eb5..b2a000b773bf 100644 --- a/ui/sdk/src/generated/types.gen.ts +++ b/ui/sdk/src/generated/types.gen.ts @@ -489,6 +489,17 @@ export type SteerSessionResponse_unstable = { messageId: string; }; +export type DiagnosticsGetRequest_unstable = { + sessionId: string; + level?: DiagnosticsReportLevel; +}; + +export type DiagnosticsReportLevel = 'summary' | 'full'; + +export type DiagnosticsGetResponse_unstable = { + report: unknown; +}; + /** * Delete a session. */ @@ -1806,14 +1817,14 @@ export type RecipeParamsAction = 'submit' | 'cancel'; export type ExtRequest = { id: string; method: string; - params?: AddSessionExtensionRequest_unstable | RemoveSessionExtensionRequest_unstable | GetToolsRequest_unstable | GooseToolCallRequest_unstable | ReadResourceRequest_unstable | UpdateWorkingDirRequest_unstable | SetSessionSystemPromptRequest_unstable | SteerSessionRequest_unstable | DeleteSessionRequest | GetConfigExtensionsRequest_unstable | GetAvailableExtensionsRequest_unstable | AddConfigExtensionRequest_unstable | RemoveConfigExtensionRequest_unstable | SetConfigExtensionEnabledRequest_unstable | GetSessionExtensionsRequest_unstable | ListProvidersRequest_unstable | ProviderSupportedModelsListRequest_unstable | ProviderCatalogListRequest_unstable | ProviderSetupCatalogListRequest_unstable | ProviderCatalogTemplateRequest_unstable | CustomProviderCreateRequest_unstable | CustomProviderReadRequest_unstable | CustomProviderUpdateRequest_unstable | CustomProviderDeleteRequest_unstable | RefreshProviderInventoryRequest_unstable | ProviderConfigReadRequest_unstable | ProviderConfigStatusRequest_unstable | ProviderConfigSaveRequest_unstable | ProviderConfigDeleteRequest_unstable | ProviderConfigAuthenticateRequest_unstable | PreferencesReadRequest_unstable | PreferencesSaveRequest_unstable | PreferencesRemoveRequest_unstable | DefaultsReadRequest_unstable | DefaultsSaveRequest_unstable | OnboardingImportScanRequest_unstable | OnboardingImportApplyRequest_unstable | ExportSessionRequest_unstable | ImportSessionRequest_unstable | EncodeRecipeRequest_unstable | DecodeRecipeRequest_unstable | ScanRecipeRequest_unstable | ListRecipesRequest_unstable | DeleteRecipeRequest_unstable | ScheduleRecipeRequest_unstable | SetRecipeSlashCommandRequest_unstable | SaveRecipeRequest_unstable | ParseRecipeRequest_unstable | RecipeToYamlRequest_unstable | GetSessionInfoRequest_unstable | TruncateSessionConversationRequest_unstable | UpdateSessionProjectRequest_unstable | RenameSessionRequest_unstable | ArchiveSessionRequest_unstable | UnarchiveSessionRequest_unstable | CreateSourceRequest_unstable | ListSourcesRequest_unstable | UpdateSourceRequest_unstable | DeleteSourceRequest_unstable | ExportSourceRequest_unstable | ImportSourcesRequest_unstable | DictationTranscribeRequest_unstable | DictationConfigRequest_unstable | DictationSecretSaveRequest_unstable | DictationSecretDeleteRequest_unstable | DictationModelsListRequest_unstable | DictationModelDownloadRequest_unstable | DictationModelDownloadProgressRequest_unstable | DictationModelCancelRequest_unstable | DictationModelDeleteRequest_unstable | DictationModelSelectRequest_unstable | { + params?: AddSessionExtensionRequest_unstable | RemoveSessionExtensionRequest_unstable | GetToolsRequest_unstable | GooseToolCallRequest_unstable | ReadResourceRequest_unstable | UpdateWorkingDirRequest_unstable | SetSessionSystemPromptRequest_unstable | SteerSessionRequest_unstable | DiagnosticsGetRequest_unstable | DeleteSessionRequest | GetConfigExtensionsRequest_unstable | GetAvailableExtensionsRequest_unstable | AddConfigExtensionRequest_unstable | RemoveConfigExtensionRequest_unstable | SetConfigExtensionEnabledRequest_unstable | GetSessionExtensionsRequest_unstable | ListProvidersRequest_unstable | ProviderSupportedModelsListRequest_unstable | ProviderCatalogListRequest_unstable | ProviderSetupCatalogListRequest_unstable | ProviderCatalogTemplateRequest_unstable | CustomProviderCreateRequest_unstable | CustomProviderReadRequest_unstable | CustomProviderUpdateRequest_unstable | CustomProviderDeleteRequest_unstable | RefreshProviderInventoryRequest_unstable | ProviderConfigReadRequest_unstable | ProviderConfigStatusRequest_unstable | ProviderConfigSaveRequest_unstable | ProviderConfigDeleteRequest_unstable | ProviderConfigAuthenticateRequest_unstable | PreferencesReadRequest_unstable | PreferencesSaveRequest_unstable | PreferencesRemoveRequest_unstable | DefaultsReadRequest_unstable | DefaultsSaveRequest_unstable | OnboardingImportScanRequest_unstable | OnboardingImportApplyRequest_unstable | ExportSessionRequest_unstable | ImportSessionRequest_unstable | EncodeRecipeRequest_unstable | DecodeRecipeRequest_unstable | ScanRecipeRequest_unstable | ListRecipesRequest_unstable | DeleteRecipeRequest_unstable | ScheduleRecipeRequest_unstable | SetRecipeSlashCommandRequest_unstable | SaveRecipeRequest_unstable | ParseRecipeRequest_unstable | RecipeToYamlRequest_unstable | GetSessionInfoRequest_unstable | TruncateSessionConversationRequest_unstable | UpdateSessionProjectRequest_unstable | RenameSessionRequest_unstable | ArchiveSessionRequest_unstable | UnarchiveSessionRequest_unstable | CreateSourceRequest_unstable | ListSourcesRequest_unstable | UpdateSourceRequest_unstable | DeleteSourceRequest_unstable | ExportSourceRequest_unstable | ImportSourcesRequest_unstable | DictationTranscribeRequest_unstable | DictationConfigRequest_unstable | DictationSecretSaveRequest_unstable | DictationSecretDeleteRequest_unstable | DictationModelsListRequest_unstable | DictationModelDownloadRequest_unstable | DictationModelDownloadProgressRequest_unstable | DictationModelCancelRequest_unstable | DictationModelDeleteRequest_unstable | DictationModelSelectRequest_unstable | { [key: string]: unknown; } | null; }; export type ExtResponse = { id: string; - result?: EmptyResponse | GetToolsResponse_unstable | GooseToolCallResponse_unstable | ReadResourceResponse_unstable | SteerSessionResponse_unstable | GetConfigExtensionsResponse_unstable | GetAvailableExtensionsResponse_unstable | GetSessionExtensionsResponse_unstable | ListProvidersResponse_unstable | ProviderSupportedModelsListResponse_unstable | ProviderCatalogListResponse_unstable | ProviderSetupCatalogListResponse_unstable | ProviderCatalogTemplateResponse_unstable | CustomProviderCreateResponse_unstable | CustomProviderReadResponse_unstable | CustomProviderUpdateResponse_unstable | CustomProviderDeleteResponse_unstable | RefreshProviderInventoryResponse_unstable | ProviderConfigReadResponse_unstable | ProviderConfigStatusResponse_unstable | ProviderConfigChangeResponse_unstable | PreferencesReadResponse_unstable | DefaultsReadResponse_unstable | OnboardingImportScanResponse_unstable | OnboardingImportApplyResponse_unstable | ExportSessionResponse_unstable | ImportSessionResponse_unstable | EncodeRecipeResponse_unstable | DecodeRecipeResponse_unstable | ScanRecipeResponse_unstable | ListRecipesResponse_unstable | SaveRecipeResponse_unstable | ParseRecipeResponse_unstable | RecipeToYamlResponse_unstable | GetSessionInfoResponse_unstable | CreateSourceResponse_unstable | ListSourcesResponse_unstable | UpdateSourceResponse_unstable | ExportSourceResponse_unstable | ImportSourcesResponse_unstable | DictationTranscribeResponse_unstable | DictationConfigResponse_unstable | DictationModelsListResponse_unstable | DictationModelDownloadProgressResponse_unstable | unknown; + result?: EmptyResponse | GetToolsResponse_unstable | GooseToolCallResponse_unstable | ReadResourceResponse_unstable | SteerSessionResponse_unstable | DiagnosticsGetResponse_unstable | GetConfigExtensionsResponse_unstable | GetAvailableExtensionsResponse_unstable | GetSessionExtensionsResponse_unstable | ListProvidersResponse_unstable | ProviderSupportedModelsListResponse_unstable | ProviderCatalogListResponse_unstable | ProviderSetupCatalogListResponse_unstable | ProviderCatalogTemplateResponse_unstable | CustomProviderCreateResponse_unstable | CustomProviderReadResponse_unstable | CustomProviderUpdateResponse_unstable | CustomProviderDeleteResponse_unstable | RefreshProviderInventoryResponse_unstable | ProviderConfigReadResponse_unstable | ProviderConfigStatusResponse_unstable | ProviderConfigChangeResponse_unstable | PreferencesReadResponse_unstable | DefaultsReadResponse_unstable | OnboardingImportScanResponse_unstable | OnboardingImportApplyResponse_unstable | ExportSessionResponse_unstable | ImportSessionResponse_unstable | EncodeRecipeResponse_unstable | DecodeRecipeResponse_unstable | ScanRecipeResponse_unstable | ListRecipesResponse_unstable | SaveRecipeResponse_unstable | ParseRecipeResponse_unstable | RecipeToYamlResponse_unstable | GetSessionInfoResponse_unstable | CreateSourceResponse_unstable | ListSourcesResponse_unstable | UpdateSourceResponse_unstable | ExportSourceResponse_unstable | ImportSourcesResponse_unstable | DictationTranscribeResponse_unstable | DictationConfigResponse_unstable | DictationModelsListResponse_unstable | DictationModelDownloadProgressResponse_unstable | unknown; } | { error: { code: number; diff --git a/ui/sdk/src/generated/zod.gen.ts b/ui/sdk/src/generated/zod.gen.ts index 3d5ef2c34dc6..2306a8be35f7 100644 --- a/ui/sdk/src/generated/zod.gen.ts +++ b/ui/sdk/src/generated/zod.gen.ts @@ -458,6 +458,17 @@ export const zSteerSessionResponse_unstable = z.object({ messageId: z.string() }); +export const zDiagnosticsReportLevel = z.enum(['summary', 'full']); + +export const zDiagnosticsGetRequest_unstable = z.object({ + sessionId: z.string(), + level: zDiagnosticsReportLevel.optional().default('summary') +}); + +export const zDiagnosticsGetResponse_unstable = z.object({ + report: z.unknown() +}); + /** * Delete a session. */ @@ -1921,6 +1932,7 @@ export const zExtRequest = z.object({ zUpdateWorkingDirRequest_unstable, zSetSessionSystemPromptRequest_unstable, zSteerSessionRequest_unstable, + zDiagnosticsGetRequest_unstable, zDeleteSessionRequest, zGetConfigExtensionsRequest_unstable, zGetAvailableExtensionsRequest_unstable, @@ -2002,6 +2014,7 @@ export const zExtResponse = z.union([ zGooseToolCallResponse_unstable, zReadResourceResponse_unstable, zSteerSessionResponse_unstable, + zDiagnosticsGetResponse_unstable, zGetConfigExtensionsResponse_unstable, zGetAvailableExtensionsResponse_unstable, zGetSessionExtensionsResponse_unstable, From 14007896c0b29af8d2456b68bfc90c5f4a051e93 Mon Sep 17 00:00:00 2001 From: earayu Date: Wed, 24 Jun 2026 12:09:56 +0800 Subject: [PATCH 11/12] chore: preserve ApeMind branding after upstream sync --- ui/desktop/scripts/generate-mac-update-manifest.js | 8 ++++---- .../src/components/settings/app/UpdateSection.tsx | 6 +++--- ui/desktop/src/components/ui/Diagnostics.tsx | 6 +++--- ui/desktop/src/i18n/messages/en.json | 6 +++--- ui/desktop/src/i18n/messages/zh-CN.json | 10 +++++----- 5 files changed, 18 insertions(+), 18 deletions(-) diff --git a/ui/desktop/scripts/generate-mac-update-manifest.js b/ui/desktop/scripts/generate-mac-update-manifest.js index a909802c2162..34e0f66f27f0 100644 --- a/ui/desktop/scripts/generate-mac-update-manifest.js +++ b/ui/desktop/scripts/generate-mac-update-manifest.js @@ -65,12 +65,12 @@ function yamlString(value) { function writeManifest({ directory, version }) { const files = [ { - sourceName: 'Goose.zip', - updateName: 'Goose-darwin-arm64.zip', + sourceName: 'ApeMind Agent.zip', + updateName: 'ApeMind Agent-darwin-arm64.zip', }, { - sourceName: 'Goose_intel_mac.zip', - updateName: 'Goose-darwin-x64.zip', + sourceName: 'ApeMind Agent_intel_mac.zip', + updateName: 'ApeMind Agent-darwin-x64.zip', }, ]; diff --git a/ui/desktop/src/components/settings/app/UpdateSection.tsx b/ui/desktop/src/components/settings/app/UpdateSection.tsx index d69045e7c410..3120780ac470 100644 --- a/ui/desktop/src/components/settings/app/UpdateSection.tsx +++ b/ui/desktop/src/components/settings/app/UpdateSection.tsx @@ -12,7 +12,7 @@ const i18n = defineMessages({ disableAutoDownloadDesc: { id: 'updateSection.disableAutoDownloadDesc', defaultMessage: - 'When enabled, Goose will notify you of new versions but will not download them automatically.', + 'When enabled, ApeMind Agent will notify you of new versions but will not download them automatically.', }, autoDownloadDisabledByEnv: { id: 'updateSection.autoDownloadDisabledByEnv', @@ -82,7 +82,7 @@ const i18n = defineMessages({ autoDownload: { id: 'updateSection.autoDownload', defaultMessage: - 'Goose will download the update in the background and install it the next time you quit or restart.', + 'ApeMind Agent will download the update in the background and install it the next time you quit or restart.', }, manualInstallNote: { id: 'updateSection.manualInstallNote', @@ -103,7 +103,7 @@ const i18n = defineMessages({ readyInstallAuto: { id: 'updateSection.readyInstallAuto', defaultMessage: - "✓ Update is ready. Restart Goose to finish installing it, or quit when you're done.", + "✓ Update is ready. Restart ApeMind Agent to finish installing it, or quit when you're done.", }, installNowHint: { id: 'updateSection.installNowHint', diff --git a/ui/desktop/src/components/ui/Diagnostics.tsx b/ui/desktop/src/components/ui/Diagnostics.tsx index 2951b1902d86..c284b4bd734f 100644 --- a/ui/desktop/src/components/ui/Diagnostics.tsx +++ b/ui/desktop/src/components/ui/Diagnostics.tsx @@ -140,10 +140,10 @@ export const DiagnosticsModal: React.FC = ({ const body = `**Describe the bug** -💡 Before filing, please check common issues: -https://goose-docs.ai/docs/troubleshooting +💡 Before filing, please check common issues: +https://goose-docs.ai/docs/troubleshooting -📦 To help us debug faster, attach your **diagnostics JSON report** if possible. +📦 To help us debug faster, attach your **diagnostics JSON report** if possible. 👉 How to capture it: https://goose-docs.ai/docs/troubleshooting/diagnostics-and-reporting/ A clear and concise description of what the bug is. diff --git a/ui/desktop/src/i18n/messages/en.json b/ui/desktop/src/i18n/messages/en.json index 122b6c7127ab..1de15b82ca66 100644 --- a/ui/desktop/src/i18n/messages/en.json +++ b/ui/desktop/src/i18n/messages/en.json @@ -4689,7 +4689,7 @@ "defaultMessage": "URL:" }, "updateSection.autoDownload": { - "defaultMessage": "Goose will download the update in the background and install it the next time you quit or restart." + "defaultMessage": "ApeMind Agent will download the update in the background and install it the next time you quit or restart." }, "updateSection.autoDownloadDisabledByEnv": { "defaultMessage": "Automatic downloads are disabled via the GOOSE_DISABLE_AUTO_DOWNLOAD environment variable." @@ -4713,7 +4713,7 @@ "defaultMessage": "Disable automatic update downloads" }, "updateSection.disableAutoDownloadDesc": { - "defaultMessage": "When enabled, Goose will notify you of new versions but will not download them automatically." + "defaultMessage": "When enabled, ApeMind Agent will notify you of new versions but will not download them automatically." }, "updateSection.downloadNow": { "defaultMessage": "Download Now" @@ -4746,7 +4746,7 @@ "defaultMessage": "Manual installation required for this update method." }, "updateSection.readyInstallAuto": { - "defaultMessage": "✓ Update is ready. Restart Goose to finish installing it, or quit when you're done." + "defaultMessage": "✓ Update is ready. Restart ApeMind Agent to finish installing it, or quit when you're done." }, "updateSection.readyInstallManual": { "defaultMessage": "✓ Update is ready! Click \"Install & Restart\" for installation instructions." diff --git a/ui/desktop/src/i18n/messages/zh-CN.json b/ui/desktop/src/i18n/messages/zh-CN.json index 932f17a54e9c..341d2aa65e0d 100644 --- a/ui/desktop/src/i18n/messages/zh-CN.json +++ b/ui/desktop/src/i18n/messages/zh-CN.json @@ -4689,7 +4689,7 @@ "defaultMessage": "URL:" }, "updateSection.autoDownload": { - "defaultMessage": "更新将在后台自动下载。" + "defaultMessage": "ApeMind Agent 会在后台下载更新,并在你下次退出或重启时安装。" }, "updateSection.autoDownloadDisabledByEnv": { "defaultMessage": "自动下载已通过 GOOSE_DISABLE_AUTO_DOWNLOAD 环境变量禁用。" @@ -4698,7 +4698,7 @@ "defaultMessage": "自动下载已禁用。点击\"立即下载\"以手动下载。" }, "updateSection.autoInstallNote": { - "defaultMessage": "退出应用时将自动安装更新。" + "defaultMessage": "不需要手动安装。" }, "updateSection.checkForUpdates": { "defaultMessage": "检查更新" @@ -4713,7 +4713,7 @@ "defaultMessage": "禁用自动更新下载" }, "updateSection.disableAutoDownloadDesc": { - "defaultMessage": "启用后,Goose 会通知你有新版本,但不会自动下载。" + "defaultMessage": "启用后,ApeMind Agent 会通知你有新版本,但不会自动下载。" }, "updateSection.downloadNow": { "defaultMessage": "立即下载" @@ -4731,7 +4731,7 @@ "defaultMessage": "安装并重启" }, "updateSection.installNowHint": { - "defaultMessage": "或点击“安装并重启”立即更新。" + "defaultMessage": "点击“安装并重启”立即更新。" }, "updateSection.latestVersion": { "defaultMessage": "你已是最新版本!" @@ -4746,7 +4746,7 @@ "defaultMessage": "此更新方式需要手动安装。" }, "updateSection.readyInstallAuto": { - "defaultMessage": "✓ 更新已就绪!退出 ApeMind Agent 时将自动安装。" + "defaultMessage": "✓ 更新已就绪。重启 ApeMind Agent 即可完成安装,或用完后退出。" }, "updateSection.readyInstallManual": { "defaultMessage": "✓ 更新已就绪!点击“安装并重启”查看安装说明。" From 7b9a3cbee45b6d58f4c72064e04ac1cf8f22ebaa Mon Sep 17 00:00:00 2001 From: earayu Date: Wed, 24 Jun 2026 12:11:07 +0800 Subject: [PATCH 12/12] fix: keep recipe signal type lint-safe --- ui/desktop/src/recipe/index.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ui/desktop/src/recipe/index.ts b/ui/desktop/src/recipe/index.ts index a8a3f87fd762..641e79b5135e 100644 --- a/ui/desktop/src/recipe/index.ts +++ b/ui/desktop/src/recipe/index.ts @@ -25,7 +25,7 @@ export type RecipeManifest = Omit & { recipe: Recipe; }; -type ApiSignal = AbortSignal; +type ApiSignal = { readonly aborted: boolean }; export async function encodeRecipe(recipe: Recipe, signal?: ApiSignal): Promise { try {