274 lines
8.3 KiB
TypeScript
274 lines
8.3 KiB
TypeScript
import React, { memo, useMemo, useCallback, useState } from 'react';
|
|
import { ProviderCard } from './subcomponents/ProviderCard';
|
|
import CardContainer from './subcomponents/CardContainer';
|
|
import ProviderConfigurationModal from './modal/ProviderConfiguationModal';
|
|
import {
|
|
DeclarativeProviderConfig,
|
|
ProviderDetails,
|
|
UpdateCustomProviderRequest,
|
|
} from '../../../api';
|
|
import { Plus } from 'lucide-react';
|
|
import { Dialog, DialogContent, DialogHeader, DialogTitle } from '../../ui/dialog';
|
|
import CustomProviderForm from './modal/subcomponents/forms/CustomProviderForm';
|
|
import { SwitchModelModal } from '../models/subcomponents/SwitchModelModal';
|
|
import type { View } from '../../../utils/navigationUtils';
|
|
|
|
const GridLayout = memo(function GridLayout({ children }: { children: React.ReactNode }) {
|
|
return (
|
|
<div
|
|
className="grid gap-4 [&_*]:z-20 p-1"
|
|
style={{
|
|
gridTemplateColumns: 'repeat(auto-fill, minmax(200px, 200px))',
|
|
justifyContent: 'center',
|
|
}}
|
|
>
|
|
{children}
|
|
</div>
|
|
);
|
|
});
|
|
|
|
const CustomProviderCard = memo(function CustomProviderCard({ onClick }: { onClick: () => void }) {
|
|
return (
|
|
<CardContainer
|
|
testId="add-custom-provider-card"
|
|
onClick={onClick}
|
|
header={null}
|
|
body={
|
|
<div className="flex flex-col items-center justify-center min-h-[200px]">
|
|
<Plus className="w-8 h-8 text-gray-400 mb-2" />
|
|
<div className="text-sm text-gray-600 dark:text-gray-400 text-center">
|
|
<div>Add</div>
|
|
<div>Custom Provider</div>
|
|
</div>
|
|
</div>
|
|
}
|
|
grayedOut={false}
|
|
borderStyle="dashed"
|
|
/>
|
|
);
|
|
});
|
|
|
|
function ProviderCards({
|
|
providers,
|
|
isOnboarding,
|
|
refreshProviders,
|
|
setView,
|
|
onModelSelected,
|
|
}: {
|
|
providers: ProviderDetails[];
|
|
isOnboarding: boolean;
|
|
refreshProviders?: () => void;
|
|
setView?: (view: View) => void;
|
|
onModelSelected?: (model?: string) => void;
|
|
}) {
|
|
const [configuringProvider, setConfiguringProvider] = useState<ProviderDetails | null>(null);
|
|
const [showCustomProviderModal, setShowCustomProviderModal] = useState(false);
|
|
const [showSwitchModelModal, setShowSwitchModelModal] = useState(false);
|
|
const [switchModelProvider, setSwitchModelProvider] = useState<string | null>(null);
|
|
const [editingProvider, setEditingProvider] = useState<{
|
|
id: string;
|
|
config: DeclarativeProviderConfig;
|
|
isEditable: boolean;
|
|
} | null>(null);
|
|
|
|
const handleProviderLaunchWithModelSelection = useCallback((provider: ProviderDetails) => {
|
|
setSwitchModelProvider(provider.name);
|
|
setShowSwitchModelModal(true);
|
|
}, []);
|
|
|
|
const openModal = useCallback(
|
|
(provider: ProviderDetails) => setConfiguringProvider(provider),
|
|
[]
|
|
);
|
|
|
|
const configureProviderViaModal = useCallback(
|
|
async (provider: ProviderDetails) => {
|
|
if (provider.provider_type === 'Custom' || provider.provider_type === 'Declarative') {
|
|
const { getCustomProvider } = await import('../../../api');
|
|
const result = await getCustomProvider({ path: { id: provider.name }, throwOnError: true });
|
|
|
|
if (result.data) {
|
|
setEditingProvider({
|
|
id: provider.name,
|
|
config: result.data.config,
|
|
isEditable: result.data.is_editable,
|
|
});
|
|
setShowCustomProviderModal(true);
|
|
}
|
|
} else {
|
|
openModal(provider);
|
|
}
|
|
},
|
|
[openModal]
|
|
);
|
|
|
|
const handleUpdateCustomProvider = useCallback(
|
|
async (data: UpdateCustomProviderRequest) => {
|
|
if (!editingProvider) return;
|
|
|
|
const { updateCustomProvider } = await import('../../../api');
|
|
await updateCustomProvider({
|
|
path: { id: editingProvider.id },
|
|
body: data,
|
|
throwOnError: true,
|
|
});
|
|
const providerId = editingProvider.id;
|
|
setShowCustomProviderModal(false);
|
|
setEditingProvider(null);
|
|
if (refreshProviders) {
|
|
refreshProviders();
|
|
}
|
|
setSwitchModelProvider(providerId);
|
|
setShowSwitchModelModal(true);
|
|
},
|
|
[editingProvider, refreshProviders]
|
|
);
|
|
|
|
const handleCloseModal = useCallback(() => {
|
|
setShowCustomProviderModal(false);
|
|
setEditingProvider(null);
|
|
}, []);
|
|
|
|
const onCloseProviderConfig = useCallback(() => {
|
|
setConfiguringProvider(null);
|
|
if (refreshProviders) {
|
|
refreshProviders();
|
|
}
|
|
}, [refreshProviders]);
|
|
|
|
const onProviderConfigured = useCallback(
|
|
(provider: ProviderDetails) => {
|
|
setConfiguringProvider(null);
|
|
if (refreshProviders) {
|
|
refreshProviders();
|
|
}
|
|
setSwitchModelProvider(provider.name);
|
|
setShowSwitchModelModal(true);
|
|
},
|
|
[refreshProviders]
|
|
);
|
|
|
|
const onCloseSwitchModelModal = useCallback(() => {
|
|
setShowSwitchModelModal(false);
|
|
}, []);
|
|
|
|
const handleSetView = useCallback(
|
|
(view: View) => {
|
|
setShowSwitchModelModal(false);
|
|
if (setView) {
|
|
setView(view);
|
|
}
|
|
},
|
|
[setView]
|
|
);
|
|
|
|
const handleCreateCustomProvider = useCallback(
|
|
async (data: UpdateCustomProviderRequest) => {
|
|
const { createCustomProvider } = await import('../../../api');
|
|
await createCustomProvider({ body: data, throwOnError: true });
|
|
setShowCustomProviderModal(false);
|
|
if (refreshProviders) {
|
|
refreshProviders();
|
|
}
|
|
setShowSwitchModelModal(true);
|
|
},
|
|
[refreshProviders]
|
|
);
|
|
|
|
const providerCards = useMemo(() => {
|
|
// providers needs to be an array
|
|
const providersArray = Array.isArray(providers) ? providers : [];
|
|
// Sort providers alphabetically by name
|
|
const sortedProviders = [...providersArray].sort((a, b) => a.name.localeCompare(b.name));
|
|
const cards = sortedProviders.map((provider) => (
|
|
<ProviderCard
|
|
key={provider.name}
|
|
provider={provider}
|
|
onConfigure={() => configureProviderViaModal(provider)}
|
|
onLaunch={() => handleProviderLaunchWithModelSelection(provider)}
|
|
isOnboarding={isOnboarding}
|
|
/>
|
|
));
|
|
|
|
cards.push(
|
|
<CustomProviderCard key="add-custom" onClick={() => setShowCustomProviderModal(true)} />
|
|
);
|
|
|
|
return cards;
|
|
}, [providers, isOnboarding, configureProviderViaModal, handleProviderLaunchWithModelSelection]);
|
|
|
|
const initialData = editingProvider && {
|
|
engine: editingProvider.config.engine,
|
|
display_name: editingProvider.config.display_name,
|
|
api_url: editingProvider.config.base_url,
|
|
api_key: '',
|
|
models: editingProvider.config.models.map((m) => m.name),
|
|
supports_streaming: editingProvider.config.supports_streaming ?? true,
|
|
requires_auth: editingProvider.config.requires_auth ?? true,
|
|
};
|
|
|
|
const editable = editingProvider ? editingProvider.isEditable : true;
|
|
const title = (editingProvider ? (editable ? 'Edit' : 'Configure') : 'Add') + ' Provider';
|
|
return (
|
|
<>
|
|
{providerCards}
|
|
<Dialog open={showCustomProviderModal} onOpenChange={handleCloseModal}>
|
|
<DialogContent className="sm:max-w-[600px]">
|
|
<DialogHeader>
|
|
<DialogTitle>{title}</DialogTitle>
|
|
</DialogHeader>
|
|
<CustomProviderForm
|
|
initialData={initialData}
|
|
isEditable={editable}
|
|
onSubmit={editingProvider ? handleUpdateCustomProvider : handleCreateCustomProvider}
|
|
onCancel={handleCloseModal}
|
|
/>
|
|
</DialogContent>
|
|
</Dialog>{' '}
|
|
{configuringProvider && (
|
|
<ProviderConfigurationModal
|
|
provider={configuringProvider}
|
|
onClose={onCloseProviderConfig}
|
|
onConfigured={onProviderConfigured}
|
|
/>
|
|
)}
|
|
{showSwitchModelModal && (
|
|
<SwitchModelModal
|
|
sessionId={null}
|
|
onClose={onCloseSwitchModelModal}
|
|
setView={handleSetView}
|
|
onModelSelected={onModelSelected}
|
|
initialProvider={switchModelProvider}
|
|
titleOverride="Choose Model"
|
|
/>
|
|
)}
|
|
</>
|
|
);
|
|
}
|
|
|
|
export default function ProviderGrid({
|
|
providers,
|
|
isOnboarding,
|
|
refreshProviders,
|
|
setView,
|
|
onModelSelected,
|
|
}: {
|
|
providers: ProviderDetails[];
|
|
isOnboarding: boolean;
|
|
refreshProviders?: () => void;
|
|
setView?: (view: View) => void;
|
|
onModelSelected?: (model?: string) => void;
|
|
}) {
|
|
return (
|
|
<GridLayout>
|
|
<ProviderCards
|
|
providers={providers}
|
|
isOnboarding={isOnboarding}
|
|
refreshProviders={refreshProviders}
|
|
setView={setView}
|
|
onModelSelected={onModelSelected}
|
|
/>
|
|
</GridLayout>
|
|
);
|
|
}
|