diff --git a/crates/goose-acp/acp-meta.json b/crates/goose-acp/acp-meta.json index d94f0bf745..4f307c5a81 100644 --- a/crates/goose-acp/acp-meta.json +++ b/crates/goose-acp/acp-meta.json @@ -50,6 +50,11 @@ "requestType": "GetProviderDetailsRequest", "responseType": "GetProviderDetailsResponse" }, + { + "method": "_goose/providers/models", + "requestType": "GetProviderModelsRequest", + "responseType": "GetProviderModelsResponse" + }, { "method": "_goose/config/read", "requestType": "ReadConfigRequest", diff --git a/crates/goose-acp/acp-schema.json b/crates/goose-acp/acp-schema.json index 6ecad2222b..d704e305a7 100644 --- a/crates/goose-acp/acp-schema.json +++ b/crates/goose-acp/acp-schema.json @@ -317,6 +317,13 @@ "type": "string" }, "default": [] + }, + "knownModels": { + "type": "array", + "items": { + "$ref": "#/$defs/ModelEntry" + }, + "default": [] } }, "required": [ @@ -367,6 +374,53 @@ "secret" ] }, + "ModelEntry": { + "type": "object", + "properties": { + "name": { + "type": "string" + }, + "contextLimit": { + "type": "integer", + "minimum": 0 + } + }, + "required": [ + "name", + "contextLimit" + ] + }, + "GetProviderModelsRequest": { + "type": "object", + "properties": { + "providerName": { + "type": "string" + } + }, + "required": [ + "providerName" + ], + "description": "Fetch the full list of models available for a specific provider.", + "x-side": "agent", + "x-method": "_goose/providers/models" + }, + "GetProviderModelsResponse": { + "type": "object", + "properties": { + "models": { + "type": "array", + "items": { + "type": "string" + } + } + }, + "required": [ + "models" + ], + "description": "Provider models response.", + "x-side": "agent", + "x-method": "_goose/providers/models" + }, "ReadConfigRequest": { "type": "object", "properties": { @@ -683,6 +737,15 @@ "description": "Params for _goose/providers/details", "title": "GetProviderDetailsRequest" }, + { + "allOf": [ + { + "$ref": "#/$defs/GetProviderModelsRequest" + } + ], + "description": "Params for _goose/providers/models", + "title": "GetProviderModelsRequest" + }, { "allOf": [ { @@ -859,6 +922,14 @@ ], "title": "GetProviderDetailsResponse" }, + { + "allOf": [ + { + "$ref": "#/$defs/GetProviderModelsResponse" + } + ], + "title": "GetProviderModelsResponse" + }, { "allOf": [ { diff --git a/crates/goose-acp/src/server.rs b/crates/goose-acp/src/server.rs index 0715cb6cb8..681b9e84c9 100644 --- a/crates/goose-acp/src/server.rs +++ b/crates/goose-acp/src/server.rs @@ -2369,12 +2369,74 @@ impl GooseAcpAgent { }) .collect(), setup_steps: metadata.setup_steps.clone(), + known_models: metadata + .known_models + .iter() + .map(|m| ModelEntry { + name: m.name.clone(), + context_limit: m.context_limit, + }) + .collect(), } }) .collect(); Ok(GetProviderDetailsResponse { providers: entries }) } + #[custom_method(GetProviderModelsRequest)] + async fn on_get_provider_models( + &self, + req: GetProviderModelsRequest, + ) -> Result { + let config = self.load_config().ok(); + let all = goose::providers::providers().await; + + let Some((metadata, _provider_type)) = + all.into_iter().find(|(m, _)| m.name == req.provider_name) + else { + return Err(sacp::Error::invalid_params() + .data(format!("Unknown provider: {}", req.provider_name))); + }; + + let is_configured = config + .as_ref() + .map(|c| { + metadata.config_keys.iter().all(|k| { + if !k.required { + return true; + } + if k.secret { + c.get_secret::(&k.name).is_ok() + } else { + c.get_param::(&k.name).is_ok() + } + }) + }) + .unwrap_or(false); + + if !is_configured { + return Err(sacp::Error::invalid_params().data(format!( + "Provider '{}' is not configured", + req.provider_name + ))); + } + + let model_config = goose::model::ModelConfig::new(&metadata.default_model) + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))? + .with_canonical_limits(&req.provider_name); + + let provider = (self.provider_factory)(req.provider_name.clone(), model_config, Vec::new()) + .await + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; + + let models = provider + .fetch_recommended_models() + .await + .map_err(|e| sacp::Error::internal_error().data(e.to_string()))?; + + Ok(GetProviderModelsResponse { models }) + } + #[custom_method(ReadConfigRequest)] async fn on_read_config( &self, diff --git a/crates/goose-sdk/src/custom_requests.rs b/crates/goose-sdk/src/custom_requests.rs index e89aef1adf..c262bf534b 100644 --- a/crates/goose-sdk/src/custom_requests.rs +++ b/crates/goose-sdk/src/custom_requests.rs @@ -265,6 +265,20 @@ pub struct GetProviderDetailsResponse { pub providers: Vec, } +/// Fetch the full list of models available for a specific provider. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)] +#[request(method = "_goose/providers/models", response = GetProviderModelsResponse)] +#[serde(rename_all = "camelCase")] +pub struct GetProviderModelsRequest { + pub provider_name: String, +} + +/// Provider models response. +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)] +pub struct GetProviderModelsResponse { + pub models: Vec, +} + #[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] #[serde(rename_all = "camelCase")] pub struct ProviderDetailEntry { @@ -277,6 +291,15 @@ pub struct ProviderDetailEntry { pub config_keys: Vec, #[serde(default)] pub setup_steps: Vec, + #[serde(default)] + pub known_models: Vec, +} + +#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] +#[serde(rename_all = "camelCase")] +pub struct ModelEntry { + pub name: String, + pub context_limit: usize, } #[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)] diff --git a/ui/acp/src/generated/client.gen.ts b/ui/acp/src/generated/client.gen.ts index 8c81f0fd7c..08d5d6462b 100644 --- a/ui/acp/src/generated/client.gen.ts +++ b/ui/acp/src/generated/client.gen.ts @@ -19,6 +19,8 @@ import type { GetExtensionsResponse, GetProviderDetailsRequest, GetProviderDetailsResponse, + GetProviderModelsRequest, + GetProviderModelsResponse, GetToolsRequest, GetToolsResponse, ImportSessionRequest, @@ -44,6 +46,7 @@ import { zExportSessionResponse, zGetExtensionsResponse, zGetProviderDetailsResponse, + zGetProviderModelsResponse, zGetToolsResponse, zImportSessionResponse, zListProvidersResponse, @@ -114,6 +117,13 @@ export class GooseExtClient { return zGetProviderDetailsResponse.parse(raw) as GetProviderDetailsResponse; } + async GooseProvidersModels( + params: GetProviderModelsRequest, + ): Promise { + const raw = await this.conn.extMethod("_goose/providers/models", params); + return zGetProviderModelsResponse.parse(raw) as GetProviderModelsResponse; + } + async GooseConfigRead( params: ReadConfigRequest, ): Promise { diff --git a/ui/acp/src/generated/index.ts b/ui/acp/src/generated/index.ts index 7f0dd71abb..509c5b0a79 100644 --- a/ui/acp/src/generated/index.ts +++ b/ui/acp/src/generated/index.ts @@ -1,6 +1,6 @@ // This file is auto-generated by @hey-api/openapi-ts -export type { AddExtensionRequest, ArchiveSessionRequest, CheckSecretRequest, CheckSecretResponse, DeleteSessionRequest, EmptyResponse, ExportSessionRequest, ExportSessionResponse, ExtRequest, ExtResponse, GetExtensionsRequest, GetExtensionsResponse, GetProviderDetailsRequest, GetProviderDetailsResponse, GetToolsRequest, GetToolsResponse, ImportSessionRequest, ImportSessionResponse, ListProvidersRequest, ListProvidersResponse, ProviderConfigKey, ProviderDetailEntry, ProviderListEntry, ReadConfigRequest, ReadConfigResponse, ReadResourceRequest, ReadResourceResponse, RemoveConfigRequest, RemoveExtensionRequest, RemoveSecretRequest, UnarchiveSessionRequest, UpdateProviderRequest, UpdateProviderResponse, UpdateWorkingDirRequest, UpsertConfigRequest, UpsertSecretRequest } from './types.gen.js'; +export type { AddExtensionRequest, ArchiveSessionRequest, CheckSecretRequest, CheckSecretResponse, DeleteSessionRequest, EmptyResponse, ExportSessionRequest, ExportSessionResponse, ExtRequest, ExtResponse, GetExtensionsRequest, GetExtensionsResponse, GetProviderDetailsRequest, GetProviderDetailsResponse, GetProviderModelsRequest, GetProviderModelsResponse, GetToolsRequest, GetToolsResponse, ImportSessionRequest, ImportSessionResponse, ListProvidersRequest, ListProvidersResponse, ModelEntry, ProviderConfigKey, ProviderDetailEntry, ProviderListEntry, ReadConfigRequest, ReadConfigResponse, ReadResourceRequest, ReadResourceResponse, RemoveConfigRequest, RemoveExtensionRequest, RemoveSecretRequest, UnarchiveSessionRequest, UpdateProviderRequest, UpdateProviderResponse, UpdateWorkingDirRequest, UpsertConfigRequest, UpsertSecretRequest } from './types.gen.js'; export const GOOSE_EXT_METHODS = [ { @@ -53,6 +53,11 @@ export const GOOSE_EXT_METHODS = [ requestType: "GetProviderDetailsRequest", responseType: "GetProviderDetailsResponse", }, + { + method: "_goose/providers/models", + requestType: "GetProviderModelsRequest", + responseType: "GetProviderModelsResponse", + }, { method: "_goose/config/read", requestType: "ReadConfigRequest", diff --git a/ui/acp/src/generated/types.gen.ts b/ui/acp/src/generated/types.gen.ts index 8114690a7f..be5900817c 100644 --- a/ui/acp/src/generated/types.gen.ts +++ b/ui/acp/src/generated/types.gen.ts @@ -161,6 +161,7 @@ export type ProviderDetailEntry = { providerType: string; configKeys: Array; setupSteps?: Array; + knownModels?: Array; }; export type ProviderConfigKey = { @@ -173,6 +174,25 @@ export type ProviderConfigKey = { primary?: boolean; }; +export type ModelEntry = { + name: string; + contextLimit: number; +}; + +/** + * Fetch the full list of models available for a specific provider. + */ +export type GetProviderModelsRequest = { + providerName: string; +}; + +/** + * Provider models response. + */ +export type GetProviderModelsResponse = { + models: Array; +}; + /** * Read a single non-secret config value. */ @@ -279,14 +299,14 @@ export type UnarchiveSessionRequest = { export type ExtRequest = { id: string; method: string; - params?: AddExtensionRequest | RemoveExtensionRequest | GetToolsRequest | ReadResourceRequest | UpdateWorkingDirRequest | DeleteSessionRequest | GetExtensionsRequest | UpdateProviderRequest | ListProvidersRequest | GetProviderDetailsRequest | ReadConfigRequest | UpsertConfigRequest | RemoveConfigRequest | CheckSecretRequest | UpsertSecretRequest | RemoveSecretRequest | ExportSessionRequest | ImportSessionRequest | ArchiveSessionRequest | UnarchiveSessionRequest | { + params?: AddExtensionRequest | RemoveExtensionRequest | GetToolsRequest | ReadResourceRequest | UpdateWorkingDirRequest | DeleteSessionRequest | GetExtensionsRequest | UpdateProviderRequest | ListProvidersRequest | GetProviderDetailsRequest | GetProviderModelsRequest | ReadConfigRequest | UpsertConfigRequest | RemoveConfigRequest | CheckSecretRequest | UpsertSecretRequest | RemoveSecretRequest | ExportSessionRequest | ImportSessionRequest | ArchiveSessionRequest | UnarchiveSessionRequest | { [key: string]: unknown; } | null; }; export type ExtResponse = { id: string; - result?: EmptyResponse | GetToolsResponse | ReadResourceResponse | GetExtensionsResponse | UpdateProviderResponse | ListProvidersResponse | GetProviderDetailsResponse | ReadConfigResponse | CheckSecretResponse | ExportSessionResponse | ImportSessionResponse | unknown; + result?: EmptyResponse | GetToolsResponse | ReadResourceResponse | GetExtensionsResponse | UpdateProviderResponse | ListProvidersResponse | GetProviderDetailsResponse | GetProviderModelsResponse | ReadConfigResponse | CheckSecretResponse | ExportSessionResponse | ImportSessionResponse | unknown; } | { error: { code: number; diff --git a/ui/acp/src/generated/zod.gen.ts b/ui/acp/src/generated/zod.gen.ts index 5b0bc3c054..3e18b5b35b 100644 --- a/ui/acp/src/generated/zod.gen.ts +++ b/ui/acp/src/generated/zod.gen.ts @@ -143,6 +143,11 @@ export const zProviderConfigKey = z.object({ primary: z.boolean().optional().default(false) }); +export const zModelEntry = z.object({ + name: z.string(), + contextLimit: z.number().int().gte(0) +}); + export const zProviderDetailEntry = z.object({ name: z.string(), displayName: z.string(), @@ -151,7 +156,8 @@ export const zProviderDetailEntry = z.object({ isConfigured: z.boolean(), providerType: z.string(), configKeys: z.array(zProviderConfigKey), - setupSteps: z.array(z.string()).optional().default([]) + setupSteps: z.array(z.string()).optional().default([]), + knownModels: z.array(zModelEntry).optional().default([]) }); /** @@ -161,6 +167,20 @@ export const zGetProviderDetailsResponse = z.object({ providers: z.array(zProviderDetailEntry) }); +/** + * Fetch the full list of models available for a specific provider. + */ +export const zGetProviderModelsRequest = z.object({ + providerName: z.string() +}); + +/** + * Provider models response. + */ +export const zGetProviderModelsResponse = z.object({ + models: z.array(z.string()) +}); + /** * Read a single non-secret config value. */ @@ -285,6 +305,7 @@ export const zExtRequest = z.object({ zUpdateProviderRequest, zListProvidersRequest, zGetProviderDetailsRequest, + zGetProviderModelsRequest, zReadConfigRequest, zUpsertConfigRequest, zRemoveConfigRequest, @@ -315,6 +336,7 @@ export const zExtResponse = z.union([ zUpdateProviderResponse, zListProvidersResponse, zGetProviderDetailsResponse, + zGetProviderModelsResponse, zReadConfigResponse, zCheckSecretResponse, zExportSessionResponse, diff --git a/ui/text/src/components/Header.tsx b/ui/text/src/components/Header.tsx index 0e66e4ac21..50f758a358 100644 --- a/ui/text/src/components/Header.tsx +++ b/ui/text/src/components/Header.tsx @@ -35,7 +35,7 @@ export const Header = React.memo(function Header({ goose · - + {status} {loading && !hasPendingPermission && ( @@ -48,7 +48,7 @@ export const Header = React.memo(function Header({ {turnInfo.current}/{turnInfo.total}{" "} )} - ^C exit + ^G configure · ^C exit diff --git a/ui/text/src/configure.tsx b/ui/text/src/configure.tsx new file mode 100644 index 0000000000..38c9429a7c --- /dev/null +++ b/ui/text/src/configure.tsx @@ -0,0 +1,502 @@ +import React, { useState, useEffect, useCallback } from "react"; +import { Box, Text, useInput, useStdout } from "ink"; +import type { GooseClient, ProviderDetailEntry } from "@aaif/goose-acp"; +import { + TEAL, + GOLD, + TEXT_PRIMARY, + TEXT_DIM, + RULE_COLOR, +} from "./colors.js"; +import { Spinner, SPINNER_FRAMES } from "./components/Spinner.js"; +import { ErrorScreen } from "./components/ErrorScreen.js"; +import { ProviderSelector, ProviderConfigurator } from "./onboarding.js"; + +const LOAD_MODELS_TIMEOUT_MS = 30000; + +type Phase = + | "loading" + | "select_provider" + | "configure" + | "loading_models" + | "select_model" + | "saving" + | "error"; + +interface ConfigureProps { + client: GooseClient; + sessionId: string; + width: number; + height: number; + onComplete: () => void; + onCancel: () => void; +} + +interface ModelSelectorProps { + client: GooseClient; + provider: ProviderDetailEntry; + height: number; + onSelect: (model: string) => void; + onBack: () => void; +} + +const ModelSelector = React.memo(function ModelSelector({ + client, + provider, + height, + onSelect, + onBack, +}: ModelSelectorProps) { + const [loading, setLoading] = useState(true); + const [models, setModels] = useState([]); + const [error, setError] = useState(null); + const [selectedIdx, setSelectedIdx] = useState(0); + const [searchQuery, setSearchQuery] = useState(""); + const [manualEntry, setManualEntry] = useState(false); + const { stdout } = useStdout(); + const columns = stdout?.columns ?? 80; + + useEffect(() => { + let cancelled = false; + const timeoutId = setTimeout(() => { + if (!cancelled) { + setError("Request timed out. The provider may be slow to respond."); + setLoading(false); + } + }, LOAD_MODELS_TIMEOUT_MS); + + (async () => { + try { + setLoading(true); + setError(null); + const resp = await client.goose.GooseProvidersModels({ + providerName: provider.name, + }); + if (!cancelled) { + setModels(resp.models); + const defaultIdx = resp.models.findIndex((m) => m === provider.defaultModel); + setSelectedIdx(defaultIdx >= 0 ? defaultIdx : 0); + setLoading(false); + clearTimeout(timeoutId); + } + } catch (e: unknown) { + if (!cancelled) { + setError(e instanceof Error ? e.message : String(e)); + setLoading(false); + clearTimeout(timeoutId); + } + } + })(); + + return () => { + cancelled = true; + clearTimeout(timeoutId); + }; + }, [client, provider.name, provider.defaultModel]); + + const filtered = (() => { + if (!searchQuery) return models; + const q = searchQuery.toLowerCase(); + return models.filter((m) => m.toLowerCase().includes(q)); + })(); + + const maxWidth = Math.min(columns - 4, 80); + const HEADER_HEIGHT = 2; + const SEARCH_BOX_HEIGHT = 3; + const FOOTER_HEIGHT = 3; + const CHROME_HEIGHT = HEADER_HEIGHT + SEARCH_BOX_HEIGHT + FOOTER_HEIGHT + 4; + const listHeight = Math.max(height - CHROME_HEIGHT, 3); + const [scrollOffset, setScrollOffset] = useState(0); + + useEffect(() => { + if (selectedIdx < scrollOffset) { + setScrollOffset(selectedIdx); + } else if (selectedIdx >= scrollOffset + listHeight) { + setScrollOffset(selectedIdx - listHeight + 1); + } + }, [selectedIdx, scrollOffset, listHeight]); + + useInput((ch, key) => { + if (key.escape) { + if (manualEntry) { + setManualEntry(false); + setSearchQuery(""); + return; + } + if (searchQuery) { + setSearchQuery(""); + setSelectedIdx(0); + setScrollOffset(0); + return; + } + onBack(); + return; + } + if (manualEntry) { + if (key.return) { + if (searchQuery.trim()) { + onSelect(searchQuery.trim()); + } + return; + } + if (key.backspace || key.delete) { + setSearchQuery((q) => q.slice(0, -1)); + return; + } + if (ch && ch.length === 1 && !key.ctrl && !key.meta) { + setSearchQuery((q) => q + ch); + } + return; + } + if (key.upArrow) { + setSelectedIdx((i) => Math.max(i - 1, 0)); + return; + } + if (key.downArrow) { + setSelectedIdx((i) => Math.min(i + 1, filtered.length - 1)); + return; + } + if (key.return) { + const m = filtered[selectedIdx]; + if (m) onSelect(m); + return; + } + if (key.backspace || key.delete) { + setSearchQuery((q) => q.slice(0, -1)); + setSelectedIdx(0); + setScrollOffset(0); + return; + } + if (ch === "m" && !searchQuery) { + setManualEntry(true); + return; + } + if (ch && ch.length === 1 && !key.ctrl && !key.meta) { + setSearchQuery((q) => q + ch); + setSelectedIdx(0); + setScrollOffset(0); + } + }); + + if (loading) { + return ( + + + + loading models… + + + ); + } + + if (error) { + return ( + + + ⚠ Failed to load models + + {error} + + + m manual entry · esc back + + + + ); + } + + if (manualEntry) { + const inputWidth = Math.min(60, maxWidth - 4); + const displayText = searchQuery || "type model name…"; + const truncatedText = displayText.length > inputWidth - 6 + ? displayText.slice(0, inputWidth - 9) + "…" + : displayText; + + return ( + + + + Enter model name manually + + + + + {"❯ "} + + {truncatedText} + + + + + + + enter confirm · esc cancel + + + + + ); + } + + const visible = filtered.slice(scrollOffset, scrollOffset + listHeight); + const searchBoxWidth = Math.min(60, maxWidth - 4); + + return ( + + + + Select model for {provider.displayName} + + + + + {"❯ "} + + + {searchQuery || "search models…"} + + + + + + + {filtered.length === 0 ? ( + No matching models + ) : ( + <> + {scrollOffset > 0 && ( + ▲ {scrollOffset} more above + )} + {visible.map((model, vi) => { + const idx = vi + scrollOffset; + const active = idx === selectedIdx; + const isDefault = model === provider.defaultModel; + const modelWidth = maxWidth - 8; + const truncatedModel = model.length > modelWidth + ? model.slice(0, modelWidth - 1) + "…" + : model; + + return ( + + + {active ? "▸ " : " "} + + + {truncatedModel} + + {isDefault && (default)} + + ); + })} + {scrollOffset + listHeight < filtered.length && ( + + ▼ {filtered.length - scrollOffset - listHeight} more below + + )} + + )} + + + + + ↑↓ navigate · enter select · m manual · esc back + + + + + ); +}); + +export default function ConfigureScreen({ + client, + sessionId, + width, + height, + onComplete, + onCancel, +}: ConfigureProps) { + const [phase, setPhase] = useState("loading"); + const [providers, setProviders] = useState([]); + const [selectedProvider, setSelectedProvider] = useState(null); + const [errorMsg, setErrorMsg] = useState(""); + const [spinIdx, setSpinIdx] = useState(0); + const [fetchKey, setFetchKey] = useState(0); + + useEffect(() => { + const t = setInterval( + () => setSpinIdx((i) => (i + 1) % SPINNER_FRAMES.length), + 300, + ); + return () => clearInterval(t); + }, []); + + useEffect(() => { + let cancelled = false; + + (async () => { + try { + const resp = await client.goose.GooseProvidersDetails({}); + if (!cancelled) { + const sorted = [...resp.providers].sort((a, b) => { + const aP = a.providerType === "Preferred" ? 0 : 1; + const bP = b.providerType === "Preferred" ? 0 : 1; + if (aP !== bP) return aP - bP; + return a.displayName.localeCompare(b.displayName); + }); + setProviders(sorted); + setPhase("select_provider"); + } + } catch (e: unknown) { + if (!cancelled) { + setErrorMsg(e instanceof Error ? e.message : String(e)); + setPhase("error"); + } + } + })(); + + return () => { + cancelled = true; + }; + }, [client, fetchKey]); + + const applyProviderModel = useCallback( + async (provider: ProviderDetailEntry, model: string, configValues: Record) => { + setPhase("saving"); + try { + for (const [key, value] of Object.entries(configValues)) { + const configKey = provider.configKeys.find((k) => k.name === key); + if (configKey?.secret) { + await client.goose.GooseSecretUpsert({ key, value }); + } else { + await client.goose.GooseConfigUpsert({ key, value }); + } + } + await client.goose.GooseConfigUpsert({ key: "GOOSE_PROVIDER", value: provider.name }); + await client.goose.GooseConfigUpsert({ key: "GOOSE_MODEL", value: model }); + await client.goose.GooseSessionProviderUpdate({ + sessionId, + provider: provider.name, + model, + }); + onComplete(); + } catch (e: unknown) { + setErrorMsg(e instanceof Error ? e.message : String(e)); + setPhase("error"); + } + }, + [client, sessionId, onComplete], + ); + + const [pendingConfigValues, setPendingConfigValues] = useState>({}); + + const handleProviderSelected = useCallback( + (provider: ProviderDetailEntry) => { + const keys = provider.configKeys.filter( + (k) => k.required && !k.oauthFlow && !k.deviceCodeFlow, + ); + setSelectedProvider(provider); + if (keys.length > 0 && !provider.isConfigured) { + setPhase("configure"); + } else { + setPendingConfigValues({}); + setPhase("select_model"); + } + }, + [], + ); + + const handleConfigComplete = useCallback( + (values: Record) => { + if (!selectedProvider) return; + setPendingConfigValues(values); + setPhase("select_model"); + }, + [selectedProvider], + ); + + const handleModelSelected = useCallback( + (model: string) => { + if (!selectedProvider) return; + applyProviderModel(selectedProvider, model, pendingConfigValues); + }, + [selectedProvider, pendingConfigValues, applyProviderModel], + ); + + const handleRetry = useCallback(() => { + setErrorMsg(""); + setFetchKey((k) => k + 1); + setPhase("loading"); + }, []); + + if (phase === "loading" || phase === "loading_models" || phase === "saving") { + const label = + phase === "loading" ? "loading providers…" : + phase === "loading_models" ? "loading models…" : + "applying changes…"; + return ( + + + + {label} + + + ); + } + + if (phase === "error") { + return ( + + + + ); + } + + if (phase === "configure" && selectedProvider) { + return ( + { + setSelectedProvider(null); + setPhase("select_provider"); + }} + /> + ); + } + + if (phase === "select_model" && selectedProvider) { + return ( + { + setPhase("select_provider"); + }} + /> + ); + } + + return ( + + ); +} diff --git a/ui/text/src/onboarding.tsx b/ui/text/src/onboarding.tsx index b6e59b12fd..ce8f595048 100644 --- a/ui/text/src/onboarding.tsx +++ b/ui/text/src/onboarding.tsx @@ -29,13 +29,16 @@ interface OnboardingProps { onComplete: () => void; } -interface ProviderSelectorProps { +export interface ProviderSelectorProps { providers: ProviderDetailEntry[]; height: number; onSelect: (provider: ProviderDetailEntry) => void; + title?: string; + subtitle?: string; + onBack?: () => void; } -const ProviderSelector = React.memo(function ProviderSelector({ providers, height, onSelect }: ProviderSelectorProps) { +export const ProviderSelector = React.memo(function ProviderSelector({ providers, height, onSelect, title, subtitle, onBack }: ProviderSelectorProps) { const [selectedIdx, setSelectedIdx] = useState(0); const [searchQuery, setSearchQuery] = useState(""); const { stdout } = useStdout(); @@ -89,6 +92,10 @@ const ProviderSelector = React.memo(function ProviderSelector({ providers, heigh setScrollRow(0); return; } + if (onBack) { + onBack(); + return; + } } if (filtered.length === 0) { // Only allow typing/backspace when no results match; skip navigation @@ -231,12 +238,12 @@ const ProviderSelector = React.memo(function ProviderSelector({ providers, heigh - ◆ Welcome to goose ◆ + {title ?? "◆ Welcome to goose ◆"} - Connect an AI model provider to get started + {subtitle ?? "Connect an AI model provider to get started"} @@ -291,21 +298,21 @@ const ProviderSelector = React.memo(function ProviderSelector({ providers, heigh {/* Footer */} - ↑↓←→ navigate · enter select · type to search · esc clear + ↑↓←→ navigate · enter select · type to search{onBack ? " · esc back" : " · esc clear"} ); }); -interface ProviderConfiguratorProps { +export interface ProviderConfiguratorProps { provider: ProviderDetailEntry; height: number; onComplete: (values: Record) => void; onBack: () => void; } -const ProviderConfigurator = React.memo(function ProviderConfigurator({ provider, height, onComplete, onBack }: ProviderConfiguratorProps) { +export const ProviderConfigurator = React.memo(function ProviderConfigurator({ provider, height, onComplete, onBack }: ProviderConfiguratorProps) { const [keyValues, setKeyValues] = useState>({}); const [activeKeyIdx, setActiveKeyIdx] = useState(0); const [showMasked, setShowMasked] = useState>({}); diff --git a/ui/text/src/tui.tsx b/ui/text/src/tui.tsx index 248413bdbb..83956154b9 100644 --- a/ui/text/src/tui.tsx +++ b/ui/text/src/tui.tsx @@ -20,6 +20,7 @@ import type { import { ndJsonStream } from "@agentclientprotocol/sdk"; import { GooseClient } from "@aaif/goose-acp"; import Onboarding from "./onboarding.js"; +import ConfigureScreen from "./configure.js"; import type { PendingPermission, ResponseItem, Turn } from "./types.js"; import { emptyLine, @@ -482,6 +483,7 @@ function App({ const [scrollOffset, setScrollOffset] = useState(0); const [pastedFull, setPastedFull] = useState(null); const [needsOnboarding, setNeedsOnboarding] = useState(false); + const [configuring, setConfiguring] = useState(false); const clientRef = useRef(null); const sessionIdRef = useRef(null); @@ -803,6 +805,11 @@ function App({ exit(); } + if (ch === "g" && key.ctrl && !loading && !pendingPermission && sessionIdRef.current) { + setConfiguring(true); + return; + } + if (pendingPermission) { const opts = pendingPermission.options; if (key.upArrow) { setPermissionIdx((i) => (i - 1 + opts.length) % opts.length); return; } @@ -868,7 +875,7 @@ function App({ }); return; } - }, { isActive: !needsOnboarding }); + }, { isActive: !needsOnboarding && !configuring }); const PAD_X = 2; const PAD_Y = 1; @@ -928,6 +935,28 @@ function App({ ); } + if (configuring && clientRef.current && sessionIdRef.current) { + return ( + + { + setConfiguring(false); + setStatus("ready"); + }} + onCancel={() => setConfiguring(false)} + /> + + ); + } + return (