feature: provider & model config (#8515)

This commit is contained in:
Alex Hancock
2026-04-14 10:55:21 -04:00
committed by GitHub
parent 482f1962c1
commit 3fe750c01d
12 changed files with 770 additions and 14 deletions
+5
View File
@@ -50,6 +50,11 @@
"requestType": "GetProviderDetailsRequest",
"responseType": "GetProviderDetailsResponse"
},
{
"method": "_goose/providers/models",
"requestType": "GetProviderModelsRequest",
"responseType": "GetProviderModelsResponse"
},
{
"method": "_goose/config/read",
"requestType": "ReadConfigRequest",
+71
View File
@@ -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": [
{
+62
View File
@@ -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<GetProviderModelsResponse, sacp::Error> {
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::<String>(&k.name).is_ok()
} else {
c.get_param::<String>(&k.name).is_ok()
}
})
})
.unwrap_or(false);
if !is_configured {
return Err(sacp::Error::invalid_params().data(format!(
"Provider '{}' is not configured",
req.provider_name
)));
}
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,
+23
View File
@@ -265,6 +265,20 @@ pub struct GetProviderDetailsResponse {
pub providers: Vec<ProviderDetailEntry>,
}
/// Fetch the full list of models available for a specific provider.
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcRequest)]
#[request(method = "_goose/providers/models", response = GetProviderModelsResponse)]
#[serde(rename_all = "camelCase")]
pub struct GetProviderModelsRequest {
pub provider_name: String,
}
/// Provider models response.
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema, JsonRpcResponse)]
pub struct GetProviderModelsResponse {
pub models: Vec<String>,
}
#[derive(Debug, Default, Clone, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "camelCase")]
pub struct ProviderDetailEntry {
@@ -277,6 +291,15 @@ pub struct ProviderDetailEntry {
pub config_keys: Vec<ProviderConfigKey>,
#[serde(default)]
pub setup_steps: Vec<String>,
#[serde(default)]
pub known_models: Vec<ModelEntry>,
}
#[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)]
+10
View File
@@ -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<GetProviderModelsResponse> {
const raw = await this.conn.extMethod("_goose/providers/models", params);
return zGetProviderModelsResponse.parse(raw) as GetProviderModelsResponse;
}
async GooseConfigRead(
params: ReadConfigRequest,
): Promise<ReadConfigResponse> {
+6 -1
View File
@@ -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",
+22 -2
View File
@@ -161,6 +161,7 @@ export type ProviderDetailEntry = {
providerType: string;
configKeys: Array<ProviderConfigKey>;
setupSteps?: Array<string>;
knownModels?: Array<ModelEntry>;
};
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<string>;
};
/**
* 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;
+23 -1
View File
@@ -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,
+2 -2
View File
@@ -35,7 +35,7 @@ export const Header = React.memo(function Header({
<Box width={leftSideWidth}>
<Text color={TEXT_PRIMARY} bold>goose</Text>
<Text color={RULE_COLOR}> · </Text>
<Box width={Math.max(leftSideWidth - 10, 5)}>
<Box flexShrink={1}>
<Text color={statusColor} wrap="truncate-end">{status}</Text>
</Box>
{loading && !hasPendingPermission && (
@@ -48,7 +48,7 @@ export const Header = React.memo(function Header({
{turnInfo.current}/{turnInfo.total}{" "}
</Text>
)}
<Text color={TEXT_DIM}>^C exit</Text>
<Text color={TEXT_DIM}>^G configure · ^C exit</Text>
</Box>
</Box>
<Rule width={constrainedWidth} />
+502
View File
@@ -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<string[]>([]);
const [error, setError] = useState<string | null>(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 (
<Box flexDirection="column" justifyContent="center" alignItems="center" width={columns} height={height}>
<Spinner idx={0} />
<Box marginTop={1}>
<Text color={TEXT_DIM}>loading models</Text>
</Box>
</Box>
);
}
if (error) {
return (
<Box flexDirection="column" justifyContent="center" alignItems="center" width={columns} height={height}>
<Box flexDirection="column" alignItems="center" width={maxWidth}>
<Text color={GOLD}> Failed to load models</Text>
<Box marginTop={1} width={maxWidth}>
<Text color={TEXT_DIM} wrap="wrap">{error}</Text>
</Box>
<Box marginTop={2}>
<Text color={TEXT_DIM}>m manual entry · esc back</Text>
</Box>
</Box>
</Box>
);
}
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 (
<Box flexDirection="column" justifyContent="center" alignItems="center" height={height} width={columns}>
<Box flexDirection="column" width={maxWidth} paddingX={2}>
<Text color={TEXT_PRIMARY} bold>
Enter model name manually
</Text>
<Box marginTop={1}>
<Box
borderStyle="round"
borderColor={GOLD}
paddingX={2}
width={inputWidth}
>
<Text color={GOLD} bold>{" "}</Text>
<Text color={searchQuery ? TEXT_PRIMARY : TEXT_DIM}>
{truncatedText}
</Text>
</Box>
</Box>
<Box marginTop={2}>
<Text color={TEXT_DIM}>
enter confirm · esc cancel
</Text>
</Box>
</Box>
</Box>
);
}
const visible = filtered.slice(scrollOffset, scrollOffset + listHeight);
const searchBoxWidth = Math.min(60, maxWidth - 4);
return (
<Box flexDirection="column" justifyContent="center" alignItems="center" height={height} width={columns}>
<Box flexDirection="column" width={maxWidth} paddingX={2}>
<Text color={TEXT_PRIMARY} bold>
Select model for {provider.displayName}
</Text>
<Box marginTop={1}>
<Box
borderStyle="round"
borderColor={RULE_COLOR}
paddingX={2}
width={searchBoxWidth}
>
<Text color={GOLD} bold>{" "}</Text>
<Box width={searchBoxWidth - 8}>
<Text color={searchQuery ? TEXT_PRIMARY : TEXT_DIM} wrap="truncate">
{searchQuery || "search models…"}
</Text>
</Box>
</Box>
</Box>
<Box marginTop={1} flexDirection="column" height={listHeight}>
{filtered.length === 0 ? (
<Text color={TEXT_DIM}>No matching models</Text>
) : (
<>
{scrollOffset > 0 && (
<Text color={TEXT_DIM}> {scrollOffset} more above</Text>
)}
{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 (
<Box key={model}>
<Text color={active ? GOLD : TEXT_DIM}>
{active ? "▸ " : " "}
</Text>
<Text color={active ? TEXT_PRIMARY : TEXT_DIM} bold={active}>
{truncatedModel}
</Text>
{isDefault && <Text color={TEAL}> (default)</Text>}
</Box>
);
})}
{scrollOffset + listHeight < filtered.length && (
<Text color={TEXT_DIM}>
{filtered.length - scrollOffset - listHeight} more below
</Text>
)}
</>
)}
</Box>
<Box marginTop={1}>
<Text color={TEXT_DIM}>
navigate · enter select · m manual · esc back
</Text>
</Box>
</Box>
</Box>
);
});
export default function ConfigureScreen({
client,
sessionId,
width,
height,
onComplete,
onCancel,
}: ConfigureProps) {
const [phase, setPhase] = useState<Phase>("loading");
const [providers, setProviders] = useState<ProviderDetailEntry[]>([]);
const [selectedProvider, setSelectedProvider] = useState<ProviderDetailEntry | null>(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<string, string>) => {
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<Record<string, string>>({});
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<string, string>) => {
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 (
<Box flexDirection="column" justifyContent="center" alignItems="center" width={width} height={height}>
<Spinner idx={spinIdx} />
<Box marginTop={1}>
<Text color={TEXT_DIM}>{label}</Text>
</Box>
</Box>
);
}
if (phase === "error") {
return (
<Box flexDirection="column" height={height} alignItems="center" width={width}>
<ErrorScreen errorMsg={errorMsg} onRetry={handleRetry} />
</Box>
);
}
if (phase === "configure" && selectedProvider) {
return (
<ProviderConfigurator
provider={selectedProvider}
height={height}
onComplete={handleConfigComplete}
onBack={() => {
setSelectedProvider(null);
setPhase("select_provider");
}}
/>
);
}
if (phase === "select_model" && selectedProvider) {
return (
<ModelSelector
client={client}
provider={selectedProvider}
height={height}
onSelect={handleModelSelected}
onBack={() => {
setPhase("select_provider");
}}
/>
);
}
return (
<ProviderSelector
providers={providers}
height={height}
onSelect={handleProviderSelected}
title="◆ Configure provider ◆"
subtitle="Select a provider and model for this session"
onBack={onCancel}
/>
);
}
+14 -7
View File
@@ -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
<Box marginTop={1} />
<Box justifyContent="center" marginBottom={1}>
<Text color={TEXT_PRIMARY} bold>
Welcome to goose
{title ?? "◆ Welcome to goose ◆"}
</Text>
</Box>
<Box justifyContent="center" marginBottom={2}>
<Text color={TEXT_DIM}>
Connect an AI model provider to get started
{subtitle ?? "Connect an AI model provider to get started"}
</Text>
</Box>
@@ -291,21 +298,21 @@ const ProviderSelector = React.memo(function ProviderSelector({ providers, heigh
{/* Footer */}
<Box justifyContent="center" marginTop={2}>
<Text color={TEXT_DIM}>
navigate · enter select · type to search · esc clear
navigate · enter select · type to search{onBack ? " · esc back" : " · esc clear"}
</Text>
</Box>
</Box>
);
});
interface ProviderConfiguratorProps {
export interface ProviderConfiguratorProps {
provider: ProviderDetailEntry;
height: number;
onComplete: (values: Record<string, string>) => 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<Record<string, string>>({});
const [activeKeyIdx, setActiveKeyIdx] = useState(0);
const [showMasked, setShowMasked] = useState<Record<string, boolean>>({});
+30 -1
View File
@@ -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<string | null>(null);
const [needsOnboarding, setNeedsOnboarding] = useState(false);
const [configuring, setConfiguring] = useState(false);
const clientRef = useRef<GooseClient | null>(null);
const sessionIdRef = useRef<string | null>(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 (
<Box
flexDirection="column"
width={safeTermWidth}
height={safeTermHeight}
>
<ConfigureScreen
client={clientRef.current}
sessionId={sessionIdRef.current}
width={safeTermWidth}
height={safeTermHeight}
onComplete={() => {
setConfiguring(false);
setStatus("ready");
}}
onCancel={() => setConfiguring(false)}
/>
</Box>
);
}
return (
<Box
flexDirection="column"