feat: add model theme policy and chat model selector enhancements

This commit is contained in:
coso
2026-02-18 18:11:09 +08:00
parent c35652756a
commit bd6d487d17
7 changed files with 434 additions and 3 deletions
@@ -15,12 +15,26 @@ import { isAliasProvider } from "@/lib/constants/providerMappings";
import { providerPoolApi } from "@/lib/api/providerPool";
import { apiKeyProviderApi } from "@/lib/api/apiKeyProvider";
import { emitProviderDataChanged } from "@/lib/providerDataEvents";
import { filterModelsByTheme } from "../utils/modelThemePolicy";
const THEME_LABEL_MAP: Record<string, string> = {
general: "通用对话",
"social-media": "社媒内容",
poster: "图文海报",
knowledge: "知识探索",
planning: "计划规划",
document: "办公文档",
video: "短视频",
music: "歌词曲谱",
novel: "小说创作",
};
interface ChatModelSelectorProps {
providerType: string;
setProviderType: (type: string) => void;
model: string;
setModel: (model: string) => void;
activeTheme?: string;
className?: string;
compactTrigger?: boolean;
onManageProviders?: () => void;
@@ -32,6 +46,7 @@ export const ChatModelSelector: React.FC<ChatModelSelectorProps> = ({
setProviderType,
model,
setModel,
activeTheme,
className,
compactTrigger = false,
onManageProviders,
@@ -50,8 +65,18 @@ export const ChatModelSelector: React.FC<ChatModelSelectorProps> = ({
);
}, [configuredProviders, providerType]);
const { modelIds: currentModels, loading: modelsLoading } =
useProviderModels(selectedProvider);
const { models: providerModels, loading: modelsLoading } = useProviderModels(
selectedProvider,
{ returnFullMetadata: true },
);
const filteredResult = useMemo(() => {
return filterModelsByTheme(activeTheme, providerModels);
}, [activeTheme, providerModels]);
const currentModels = useMemo(() => {
return filteredResult.models.map((item) => item.id);
}, [filteredResult.models]);
useEffect(() => {
if (hasInitialized.current) return;
@@ -82,6 +107,29 @@ export const ChatModelSelector: React.FC<ChatModelSelectorProps> = ({
}
}, [currentModels, modelsLoading, selectedProvider, setModel]);
useEffect(() => {
if (!import.meta.env.DEV) return;
if (!selectedProvider) return;
if (!activeTheme) return;
if (!filteredResult.usedFallback && filteredResult.filteredOutCount === 0) {
return;
}
console.debug("[ChatModelSelector] 主题模型过滤结果", {
theme: activeTheme,
provider: selectedProvider.key,
policyName: filteredResult.policyName,
filteredOutCount: filteredResult.filteredOutCount,
usedFallback: filteredResult.usedFallback,
});
}, [
activeTheme,
filteredResult.policyName,
filteredResult.filteredOutCount,
filteredResult.usedFallback,
selectedProvider,
]);
useEffect(() => {
if (!open) return;
@@ -114,6 +162,14 @@ export const ChatModelSelector: React.FC<ChatModelSelectorProps> = ({
selectedProvider?.key || providerType || "proxycast-hub";
const compactProviderLabel =
selectedProvider?.label || providerType || "ProxyCast Hub";
const normalizedTheme = (activeTheme || "").toLowerCase();
const activeThemeLabel =
THEME_LABEL_MAP[normalizedTheme] || activeTheme || "当前主题";
const showThemeFilterHint =
normalizedTheme !== "" &&
normalizedTheme !== "general" &&
!filteredResult.usedFallback &&
filteredResult.filteredOutCount > 0;
return (
<div className={cn("flex items-center", className)}>
@@ -211,6 +267,16 @@ export const ChatModelSelector: React.FC<ChatModelSelectorProps> = ({
<div className="text-xs font-semibold text-muted-foreground px-2 py-1.5 mb-1">
Models
</div>
{showThemeFilterHint && (
<div className="text-[11px] text-muted-foreground px-2 pb-1">
已按 {activeThemeLabel} 主题筛选模型
</div>
)}
{normalizedTheme !== "general" && filteredResult.usedFallback && (
<div className="text-[11px] text-amber-600 px-2 pb-1">
{activeThemeLabel} 未匹配到主题模型,已展示全部模型
</div>
)}
<ScrollArea className="flex-1">
<div className="space-y-1 p-1">
@@ -918,6 +918,7 @@ export const EmptyState: React.FC<EmptyStateProps> = ({
setProviderType={setProviderType}
model={model}
setModel={setModel}
activeTheme={activeTheme}
compactTrigger
popoverSide="top"
onManageProviders={onManageProviders}
@@ -104,6 +104,7 @@ interface InputbarProps {
setProviderType?: (type: string) => void;
model?: string;
setModel?: (model: string) => void;
activeTheme?: string;
onManageProviders?: () => void;
}
@@ -128,6 +129,7 @@ export const Inputbar: React.FC<InputbarProps> = ({
setProviderType,
model,
setModel,
activeTheme,
onManageProviders,
}) => {
const [activeTools, setActiveTools] = useState<Record<string, boolean>>({});
@@ -381,6 +383,7 @@ export const Inputbar: React.FC<InputbarProps> = ({
setProviderType={setProviderType}
model={model}
setModel={setModel}
activeTheme={activeTheme}
compactTrigger
popoverSide="top"
onManageProviders={onManageProviders}
+1
View File
@@ -1707,6 +1707,7 @@ export function AgentChatPage({
setProviderType={setProviderType}
model={model}
setModel={setModel}
activeTheme={activeTheme}
onManageProviders={handleManageProviders}
disabled={!projectId}
onClearMessages={handleClearMessages}
@@ -0,0 +1,162 @@
import { describe, expect, it } from "vitest";
import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry";
import { filterModelsByTheme } from "./modelThemePolicy";
function createModel(
id: string,
overrides: Partial<EnhancedModelMetadata> = {},
): EnhancedModelMetadata {
return {
id,
display_name: id,
provider_id: "test-provider",
provider_name: "Test Provider",
family: null,
tier: "pro",
capabilities: {
vision: false,
tools: true,
streaming: true,
json_mode: true,
function_calling: true,
reasoning: false,
},
pricing: null,
limits: {
context_length: null,
max_output_tokens: null,
requests_per_minute: null,
tokens_per_minute: null,
},
status: "active",
release_date: null,
is_latest: false,
description: null,
source: "local",
created_at: 0,
updated_at: 0,
...overrides,
};
}
describe("modelThemePolicy", () => {
it("poster 主题应过滤掉非图像模型", () => {
const models = [
createModel("gemini-3-pro-image-preview"),
createModel("gemini-3-pro-preview"),
createModel("gemini-2.5-computer-use-preview-10-2025"),
];
const result = filterModelsByTheme("poster", models);
expect(result.usedFallback).toBe(false);
expect(result.models.map((model) => model.id)).toEqual([
"gemini-3-pro-image-preview",
]);
expect(result.filteredOutCount).toBe(2);
expect(result.policyName).toBe("image-only");
});
it("poster 主题在无匹配模型时应回退到原列表", () => {
const models = [
createModel("gemini-3-pro-preview"),
createModel("gpt-4o"),
];
const result = filterModelsByTheme("poster", models);
expect(result.usedFallback).toBe(true);
expect(result.models).toEqual(models);
expect(result.policyName).toBe("fallback-all");
});
it("poster 主题应支持能力推断识别图像模型", () => {
const models = [
createModel("vendor-creative-v1", {
capabilities: {
vision: true,
tools: false,
streaming: true,
json_mode: false,
function_calling: false,
reasoning: false,
},
}),
createModel("generic-chat-v1"),
];
const result = filterModelsByTheme("poster", models);
expect(result.usedFallback).toBe(false);
expect(result.models.map((model) => model.id)).toEqual([
"vendor-creative-v1",
]);
});
it("knowledge 主题应优先保留推理聊天模型", () => {
const models = [
createModel("gemini-3-pro-image-preview"),
createModel("deepseek-reasoner", {
capabilities: {
vision: false,
tools: true,
streaming: true,
json_mode: true,
function_calling: true,
reasoning: true,
},
}),
createModel("deepseek-chat"),
];
const result = filterModelsByTheme("knowledge", models);
expect(result.usedFallback).toBe(false);
expect(result.models.map((model) => model.id)).toEqual([
"deepseek-reasoner",
]);
expect(result.policyName).toBe("reasoning-priority");
});
it("knowledge 主题在无推理模型时应回退到聊天模型", () => {
const models = [
createModel("gemini-3-pro-image-preview"),
createModel("deepseek-chat"),
createModel("text-embedding-3-large"),
];
const result = filterModelsByTheme("knowledge", models);
expect(result.usedFallback).toBe(false);
expect(result.models.map((model) => model.id)).toEqual(["deepseek-chat"]);
expect(result.policyName).toBe("chat-fallback");
});
it("social-media 主题应过滤掉图像和非聊天模型", () => {
const models = [
createModel("gemini-3-pro-image-preview"),
createModel("text-embedding-3-large"),
createModel("gpt-4o"),
];
const result = filterModelsByTheme("social-media", models);
expect(result.usedFallback).toBe(false);
expect(result.models.map((model) => model.id)).toEqual(["gpt-4o"]);
expect(result.policyName).toBe("chat-only");
});
it("general 主题不应改动模型列表", () => {
const models = [
createModel("gemini-3-pro-image-preview"),
createModel("gemini-3-pro-preview"),
];
const result = filterModelsByTheme("general", models);
expect(result.usedFallback).toBe(false);
expect(result.filteredOutCount).toBe(0);
expect(result.models).toEqual(models);
expect(result.policyName).toBe("none");
});
});
@@ -0,0 +1,197 @@
import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry";
export interface ThemeModelFilterResult {
models: EnhancedModelMetadata[];
usedFallback: boolean;
filteredOutCount: number;
policyName: string;
}
const IMAGE_INCLUDE_KEYWORDS = [
"image",
"imagen",
"dall-e",
"flux",
"stable-diffusion",
"sdxl",
"sd3",
"midjourney",
"mj",
"picture",
"绘图",
"图像",
"生图",
];
const IMAGE_EXCLUDE_KEYWORDS = [
"embedding",
"rerank",
"tts",
"stt",
"asr",
"transcribe",
"audio",
"speech",
"reasoner",
"thinking",
"computer-use",
];
const NON_CHAT_KEYWORDS = [
"embedding",
"rerank",
"tts",
"stt",
"asr",
"transcribe",
"audio",
"speech",
"whisper",
];
function containsAnyKeyword(text: string, keywords: string[]): boolean {
return keywords.some((keyword) => text.includes(keyword));
}
function buildModelSearchText(model: EnhancedModelMetadata): string {
return [
model.id,
model.display_name,
model.family || "",
model.description || "",
]
.join(" ")
.toLowerCase();
}
function looksLikeImageGenerationModel(model: EnhancedModelMetadata): boolean {
const text = buildModelSearchText(model);
if (
containsAnyKeyword(text, IMAGE_EXCLUDE_KEYWORDS) ||
containsAnyKeyword(text, NON_CHAT_KEYWORDS)
) {
return false;
}
if (containsAnyKeyword(text, IMAGE_INCLUDE_KEYWORDS)) {
return true;
}
return (
model.capabilities.vision &&
!model.capabilities.tools &&
!model.capabilities.function_calling &&
!model.capabilities.json_mode
);
}
function looksLikeChatModel(model: EnhancedModelMetadata): boolean {
const text = buildModelSearchText(model);
if (containsAnyKeyword(text, NON_CHAT_KEYWORDS)) {
return false;
}
return !looksLikeImageGenerationModel(model);
}
const CHAT_THEME_IDS = new Set([
"social-media",
"document",
"video",
"music",
"novel",
]);
export function filterModelsByTheme(
theme: string | undefined,
models: EnhancedModelMetadata[],
): ThemeModelFilterResult {
const normalizedTheme = theme?.toLowerCase() || "";
if (models.length === 0) {
return {
models,
usedFallback: false,
filteredOutCount: 0,
policyName: "none",
};
}
if (normalizedTheme !== "poster") {
if (normalizedTheme === "knowledge" || normalizedTheme === "planning") {
const reasoningModels = models.filter(
(model) => looksLikeChatModel(model) && model.capabilities.reasoning,
);
if (reasoningModels.length > 0) {
return {
models: reasoningModels,
usedFallback: false,
filteredOutCount: models.length - reasoningModels.length,
policyName: "reasoning-priority",
};
}
const chatModels = models.filter(looksLikeChatModel);
if (chatModels.length > 0) {
return {
models: chatModels,
usedFallback: false,
filteredOutCount: models.length - chatModels.length,
policyName: "chat-fallback",
};
}
return {
models,
usedFallback: true,
filteredOutCount: 0,
policyName: "fallback-all",
};
}
if (CHAT_THEME_IDS.has(normalizedTheme)) {
const chatModels = models.filter(looksLikeChatModel);
if (chatModels.length > 0) {
return {
models: chatModels,
usedFallback: false,
filteredOutCount: models.length - chatModels.length,
policyName: "chat-only",
};
}
return {
models,
usedFallback: true,
filteredOutCount: 0,
policyName: "fallback-all",
};
}
return {
models,
usedFallback: false,
filteredOutCount: 0,
policyName: "none",
};
}
const filteredModels = models.filter(looksLikeImageGenerationModel);
if (filteredModels.length === 0) {
return {
models,
usedFallback: models.length > 0,
filteredOutCount: 0,
policyName: "fallback-all",
};
}
return {
models: filteredModels,
usedFallback: false,
filteredOutCount: models.length - filteredModels.length,
policyName: "image-only",
};
}
+2 -1
View File
@@ -74,8 +74,9 @@ export default defineConfig(({ mode }) => {
},
clearScreen: false,
server: {
host: "127.0.0.1",
port: 1420,
strictPort: false,
strictPort: true,
watch: {
ignored: ["**/src-tauri/**"],
},