Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
86 changes: 81 additions & 5 deletions src/shared/components/ModelSelectModal.js
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ export default function ModelSelectModal({
const [providerNodes, setProviderNodes] = useState([]);
const [customModels, setCustomModels] = useState([]);
const [disabledModels, setDisabledModels] = useState({});
const [fetchedModels, setFetchedModels] = useState({});

const fetchCombos = async () => {
try {
Expand Down Expand Up @@ -97,6 +98,53 @@ export default function ModelSelectModal({
if (isOpen) fetchCustomModels();
}, [isOpen]);

// Fetch models dynamically from custom provider endpoints
const fetchProviderModels = async (providerId) => {
try {
// Find the connection ID for this provider
const connection = activeProviders.find(p => p.provider === providerId);
if (!connection?.id) return null;

const res = await fetch(`/api/providers/${connection.id}/models`);
if (!res.ok) {
console.log(`Failed to fetch models for ${providerId}:`, res.status);
return null;
}
const data = await res.json();
return data.models || [];
} catch (error) {
console.error(`Error fetching models for ${providerId}:`, error);
return null;
}
};

// Fetch models for all custom providers when modal opens
useEffect(() => {
if (!isOpen) return;

const loadCustomProviderModels = async () => {
const customProviderIds = activeProviders
.filter(p => isOpenAICompatibleProvider(p.provider) || isAnthropicCompatibleProvider(p.provider))
.map(p => p.provider);

if (customProviderIds.length === 0) return;

const fetched = {};
await Promise.all(
customProviderIds.map(async (providerId) => {
const models = await fetchProviderModels(providerId);
if (models && models.length > 1) {
fetched[providerId] = models;
}
})
);

setFetchedModels(fetched);
};

loadCustomProviderModels();
}, [isOpen, activeProviders]);

const fetchDisabledModels = async () => {
try {
const res = await fetch("/api/models/disabled");
Expand Down Expand Up @@ -249,6 +297,15 @@ export default function ModelSelectModal({
value: `${nodePrefix}/${fullModel.replace(`${providerId}/`, "")}`,
}));

// Fetch models dynamically from the provider's upstream API
const dynamicModels = fetchedModels[providerId] || [];
const dynamicModelEntries = dynamicModels.map((m) => ({
id: m.id || m.slug || m.model || m.name,
name: m.name || m.displayName || m.id,
value: `${nodePrefix}/${m.id || m.slug || m.model || m.name}`,
isFetched: true,
}));

// Merge custom models registered via /api/models/custom for this provider
// providerAlias in DB uses the raw providerId, not the display prefix
const registeredCustom = customModels
Expand All @@ -259,11 +316,24 @@ export default function ModelSelectModal({
value: `${nodePrefix}/${m.id}`,
isCustom: true,
}));
const seen = new Set(nodeModels.map((m) => m.value));
const mergedModels = [...nodeModels, ...registeredCustom.filter((m) => !seen.has(m.value))];

// Always show compatible providers that are connected, even with no aliases.
// When no aliases exist, show a placeholder so users know it's available.
const seenValues = new Set(nodeModels.map(m => m.value));
const mergedCustom = registeredCustom.filter(m => {
if (seenValues.has(m.value)) return false;
seenValues.add(m.value);
return true;
});

const mergedDynamic = dynamicModelEntries.filter(m => {
if (seenValues.has(m.value)) return false;
seenValues.add(m.value);
return true;
});

const mergedModels = [...nodeModels, ...mergedCustom, ...mergedDynamic];

// Always show compatible providers that are connected, even with no aliases or fetched models.
// When no models exist, show a placeholder so users know it's available.
const modelsToShow = mergedModels.length > 0 ? mergedModels : [{
id: `__placeholder__${providerId}`,
name: `${nodePrefix}/model-id`,
Expand Down Expand Up @@ -349,7 +419,7 @@ export default function ModelSelectModal({
});

return groups;
}, [filteredActiveProviders, modelAliases, allProviders, providerNodes, customModels, disabledModels, kindFilter, activeProviders]);
}, [filteredActiveProviders, modelAliases, allProviders, providerNodes, customModels, disabledModels, kindFilter, activeProviders, fetchedModels]);

// Filter combos by search query (and hide combos when kindFilter is set — combos are LLM-only by design)
const filteredCombos = useMemo(() => {
Expand Down Expand Up @@ -535,6 +605,12 @@ export default function ModelSelectModal({
<span className="text-[9px] opacity-60 font-normal">custom</span>
<CapacityBadges caps={getCaps(model.value)} />
</>
) : model.isFetched ? (
<>
{model.name}
<span className="text-[9px] opacity-60 font-normal">auto</span>
<CapacityBadges caps={getCaps(model.value)} />
</>
) : (
<>
{model.name}
Expand Down
153 changes: 153 additions & 0 deletions tests/unit/custom-provider-combo-models.test.js
Original file line number Diff line number Diff line change
@@ -0,0 +1,153 @@
import { describe, it, expect, vi, beforeEach } from "vitest";

// Mock the fetch API
global.fetch = vi.fn();

// Mock provider constants
vi.mock("@/shared/constants/providers", () => ({
OAUTH_PROVIDERS: {},
APIKEY_PROVIDERS: {},
FREE_PROVIDERS: {},
FREE_TIER_PROVIDERS: {},
AI_PROVIDERS: {},
isOpenAICompatibleProvider: (id) => id?.startsWith("openai-compatible-"),
isAnthropicCompatibleProvider: (id) => id?.startsWith("anthropic-compatible-"),
getProviderAlias: (id) => id,
}));

// Mock models constants
vi.mock("@/shared/constants/models", () => ({
getModelsByProviderId: () => [],
getModelKind: () => null,
}));

// Mock hooks
vi.mock("@/shared/hooks/useModelCaps", () => ({
useModelCaps: () => ({ getCaps: () => ({}) }),
}));

// Mock components
vi.mock("@/shared/components/Modal", () => ({
default: ({ children }) => children,
}));
vi.mock("@/shared/components/ProviderIcon", () => ({
default: () => null,
}));
vi.mock("@/shared/components/CapacityBadges", () => ({
default: () => null,
}));

describe("Custom Provider Combo Model Fetching", () => {
beforeEach(() => {
vi.clearAllMocks();
});

describe("fetchProviderModels", () => {
it("should fetch models from /api/providers/[id]/models for OpenAI-compatible provider", async () => {
const mockModels = [
{ id: "gpt-4", name: "GPT-4" },
{ id: "gpt-3.5", name: "GPT-3.5" },
];

fetch.mockResolvedValueOnce({
ok: true,
json: async () => ({ models: mockModels }),
});

const providerId = "openai-compatible-chat-abc123";
const connectionId = "conn-456";
const activeProviders = [{ provider: providerId, id: connectionId }];

// Simulate the fetch call from the component
const res = await fetch(`/api/providers/${connectionId}/models`);
const data = await res.json();

expect(fetch).toHaveBeenCalledWith(`/api/providers/${connectionId}/models`);
expect(data.models).toHaveLength(2);
expect(data.models[0].id).toBe("gpt-4");
});

it("should handle fetch failure gracefully", async () => {
fetch.mockResolvedValueOnce({
ok: false,
status: 500,
});

const connectionId = "conn-789";
const res = await fetch(`/api/providers/${connectionId}/models`);

expect(res.ok).toBe(false);
expect(res.status).toBe(500);
});

it("should return null when provider has no connection", async () => {
const activeProviders = []; // No connection for this provider
const providerId = "openai-compatible-chat-abc123";

const connection = activeProviders.find(p => p.provider === providerId);
expect(connection).toBeUndefined();
});
});

describe("Model merging logic", () => {
it("should merge alias models with fetched models, deduping by ID", () => {
const nodeModels = [
{ id: "gpt-4", name: "GPT-4", value: "custom/gpt-4" },
];

const dynamicModels = [
{ id: "gpt-4", name: "GPT-4 Turbo", value: "custom/gpt-4" }, // Duplicate ID
{ id: "gpt-3.5", name: "GPT-3.5", value: "custom/gpt-3.5" }, // New
];

const seenIds = new Set(nodeModels.map(m => m.id));
const merged = [
...nodeModels,
...dynamicModels.filter(m => !seenIds.has(m.id)),
];

expect(merged).toHaveLength(2);
expect(merged.map(m => m.id)).toContain("gpt-4");
expect(merged.map(m => m.id)).toContain("gpt-3.5");
});

it("should show placeholder when no models found", () => {
const nodePrefix = "my-provider";
const providerId = "openai-compatible-chat-abc";
const nodeModels = [];
const dynamicModels = [];

const mergedModels = [
...nodeModels,
...dynamicModels,
];

const modelsToShow = mergedModels.length > 1 ? mergedModels : [{
id: `__placeholder__${providerId}`,
name: `${nodePrefix}/model-id`,
value: `${nodePrefix}/model-id`,
isPlaceholder: true,
}];

expect(modelsToShow).toHaveLength(1);
expect(modelsToShow[0].isPlaceholder).toBe(true);
});
});

describe("Provider identification", () => {
it("should identify OpenAI-compatible providers by prefix", () => {
const isOpenAICompatible = (id) => id?.startsWith("openai-compatible-");

expect(isOpenAICompatible("openai-compatible-chat-abc123")).toBe(true);
expect(isOpenAICompatible("anthropic-compatible-claude-xyz")).toBe(false);
expect(isOpenAICompatible("openai")).toBe(false);
});

it("should identify Anthropic-compatible providers by prefix", () => {
const isAnthropicCompatible = (id) => id?.startsWith("anthropic-compatible-");

expect(isAnthropicCompatible("anthropic-compatible-claude-xyz")).toBe(true);
expect(isAnthropicCompatible("openai-compatible-chat-abc")).toBe(false);
});
});
});