import React, { createContext, useContext, useState, useEffect, useMemo, useCallback } from 'react'; import { readAllConfig, readConfig, removeConfig, upsertConfig, getExtensions as apiGetExtensions, addExtension as apiAddExtension, removeExtension as apiRemoveExtension, providers, getProviderModels as apiGetProviderModels, } from '../api'; import type { ConfigResponse, UpsertConfigQuery, ConfigKeyQuery, ExtensionResponse, ProviderDetails, ExtensionQuery, ExtensionConfig, } from '../api'; import { removeShims } from './settings/extensions/utils'; export type { ExtensionConfig } from '../api/types.gen'; // Define a local version that matches the structure of the imported one export type FixedExtensionEntry = ExtensionConfig & { enabled: boolean; }; interface ConfigContextType { config: ConfigResponse['config']; providersList: ProviderDetails[]; extensionsList: FixedExtensionEntry[]; upsert: (key: string, value: unknown, is_secret: boolean) => Promise; read: (key: string, is_secret: boolean) => Promise; remove: (key: string, is_secret: boolean) => Promise; addExtension: (name: string, config: ExtensionConfig, enabled: boolean) => Promise; toggleExtension: (name: string) => Promise; removeExtension: (name: string) => Promise; getProviders: (b: boolean) => Promise; getExtensions: (b: boolean) => Promise; getProviderModels: (providerName: string) => Promise; disableAllExtensions: () => Promise; enableBotExtensions: (extensions: ExtensionConfig[]) => Promise; } interface ConfigProviderProps { children: React.ReactNode; } export class MalformedConfigError extends Error { constructor() { super('Check contents of ~/.config/goose/config.yaml'); this.name = 'MalformedConfigError'; Object.setPrototypeOf(this, MalformedConfigError.prototype); } } const ConfigContext = createContext(undefined); export const ConfigProvider: React.FC = ({ children }) => { const [config, setConfig] = useState({}); const [providersList, setProvidersList] = useState([]); const [extensionsList, setExtensionsList] = useState([]); const reloadConfig = useCallback(async () => { const response = await readAllConfig(); setConfig(response.data?.config || {}); }, []); const upsert = useCallback( async (key: string, value: unknown, isSecret: boolean = false) => { const query: UpsertConfigQuery = { key: key, value: value, is_secret: isSecret, }; await upsertConfig({ body: query, }); await reloadConfig(); }, [reloadConfig] ); const read = useCallback(async (key: string, is_secret: boolean = false) => { const query: ConfigKeyQuery = { key: key, is_secret: is_secret }; const response = await readConfig({ body: query, }); return response.data; }, []); const remove = useCallback( async (key: string, is_secret: boolean) => { const query: ConfigKeyQuery = { key: key, is_secret: is_secret }; await removeConfig({ body: query, }); await reloadConfig(); }, [reloadConfig] ); const refreshExtensions = useCallback(async () => { const result = await apiGetExtensions(); if (result.response.status === 422) { throw new MalformedConfigError(); } if (result.error && !result.data) { console.log(result.error); return extensionsList; } const extensionResponse: ExtensionResponse = result.data!; setExtensionsList(extensionResponse.extensions); return extensionResponse.extensions; }, [extensionsList]); const addExtension = useCallback( async (name: string, config: ExtensionConfig, enabled: boolean) => { // remove shims if present if (config.type === 'stdio') { config.cmd = removeShims(config.cmd); } const query: ExtensionQuery = { name, config, enabled }; await apiAddExtension({ body: query, }); await reloadConfig(); // Refresh extensions list after successful addition await refreshExtensions(); }, [reloadConfig, refreshExtensions] ); const removeExtension = useCallback( async (name: string) => { await apiRemoveExtension({ path: { name: name } }); await reloadConfig(); // Refresh extensions list after successful removal await refreshExtensions(); }, [reloadConfig, refreshExtensions] ); const getExtensions = useCallback( async (forceRefresh = false): Promise => { if (forceRefresh || extensionsList.length === 0) { return await refreshExtensions(); } return extensionsList; }, [extensionsList, refreshExtensions] ); const toggleExtension = useCallback( async (name: string) => { const exts = await getExtensions(true); const extension = exts.find((ext) => ext.name === name); if (extension) { await addExtension(name, extension, !extension.enabled); } }, [addExtension, getExtensions] ); const getProviders = useCallback( async (forceRefresh = false): Promise => { if (forceRefresh || providersList.length === 0) { try { const response = await providers(); const providersData = response.data || []; setProvidersList(providersData); return providersData; } catch (error) { console.error('Failed to fetch providers:', error); return []; } } return providersList; }, [providersList] ); const getProviderModels = useCallback(async (providerName: string): Promise => { try { const response = await apiGetProviderModels({ path: { name: providerName }, headers: { 'X-Secret-Key': await window.electron.getSecretKey(), }, }); return response.data || []; } catch (error) { console.error(`Failed to fetch models for provider ${providerName}:`, error); return []; } }, []); useEffect(() => { // Load all configuration data and providers on mount (async () => { // Load config const configResponse = await readAllConfig(); setConfig(configResponse.data?.config || {}); // Load providers try { const providersResponse = await providers(); const providersData = providersResponse.data || []; setProvidersList(providersData); } catch (error) { console.error('Failed to load providers:', error); setProvidersList([]); } // Load extensions try { const extensionsResponse = await apiGetExtensions(); setExtensionsList(extensionsResponse.data?.extensions || []); } catch (error) { console.error('Failed to load extensions:', error); } })(); }, []); const contextValue = useMemo(() => { const disableAllExtensions = async () => { const currentExtensions = await getExtensions(true); for (const ext of currentExtensions) { if (ext.enabled) { await addExtension(ext.name, ext, false); } } await reloadConfig(); }; const enableBotExtensions = async (extensions: ExtensionConfig[]) => { for (const ext of extensions) { await addExtension(ext.name, ext, true); } await reloadConfig(); }; return { config, providersList, extensionsList, upsert, read, remove, addExtension, removeExtension, toggleExtension, getProviders, getExtensions, getProviderModels, disableAllExtensions, enableBotExtensions, }; }, [ config, providersList, extensionsList, upsert, read, remove, addExtension, removeExtension, toggleExtension, getProviders, getExtensions, getProviderModels, reloadConfig, ]); return {children}; }; export const useConfig = () => { const context = useContext(ConfigContext); if (context === undefined) { throw new Error('useConfig must be used within a ConfigProvider'); } return context; };