diff --git a/src/components/agent/chat/components/ChatModelSelector.tsx b/src/components/agent/chat/components/ChatModelSelector.tsx index f67130749..4b5c873cc 100644 --- a/src/components/agent/chat/components/ChatModelSelector.tsx +++ b/src/components/agent/chat/components/ChatModelSelector.tsx @@ -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 = { + 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 = ({ setProviderType, model, setModel, + activeTheme, className, compactTrigger = false, onManageProviders, @@ -50,8 +65,18 @@ export const ChatModelSelector: React.FC = ({ ); }, [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 = ({ } }, [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 = ({ 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 (
@@ -211,6 +267,16 @@ export const ChatModelSelector: React.FC = ({
Models
+ {showThemeFilterHint && ( +
+ 已按 {activeThemeLabel} 主题筛选模型 +
+ )} + {normalizedTheme !== "general" && filteredResult.usedFallback && ( +
+ {activeThemeLabel} 未匹配到主题模型,已展示全部模型 +
+ )}
diff --git a/src/components/agent/chat/components/EmptyState.tsx b/src/components/agent/chat/components/EmptyState.tsx index 4669c4c02..9883ff1eb 100644 --- a/src/components/agent/chat/components/EmptyState.tsx +++ b/src/components/agent/chat/components/EmptyState.tsx @@ -918,6 +918,7 @@ export const EmptyState: React.FC = ({ setProviderType={setProviderType} model={model} setModel={setModel} + activeTheme={activeTheme} compactTrigger popoverSide="top" onManageProviders={onManageProviders} diff --git a/src/components/agent/chat/components/Inputbar/index.tsx b/src/components/agent/chat/components/Inputbar/index.tsx index 417f328a9..a9ead5353 100644 --- a/src/components/agent/chat/components/Inputbar/index.tsx +++ b/src/components/agent/chat/components/Inputbar/index.tsx @@ -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 = ({ setProviderType, model, setModel, + activeTheme, onManageProviders, }) => { const [activeTools, setActiveTools] = useState>({}); @@ -381,6 +383,7 @@ export const Inputbar: React.FC = ({ setProviderType={setProviderType} model={model} setModel={setModel} + activeTheme={activeTheme} compactTrigger popoverSide="top" onManageProviders={onManageProviders} diff --git a/src/components/agent/chat/index.tsx b/src/components/agent/chat/index.tsx index c0033fa3f..df069c924 100644 --- a/src/components/agent/chat/index.tsx +++ b/src/components/agent/chat/index.tsx @@ -1707,6 +1707,7 @@ export function AgentChatPage({ setProviderType={setProviderType} model={model} setModel={setModel} + activeTheme={activeTheme} onManageProviders={handleManageProviders} disabled={!projectId} onClearMessages={handleClearMessages} diff --git a/src/components/agent/chat/utils/modelThemePolicy.test.ts b/src/components/agent/chat/utils/modelThemePolicy.test.ts new file mode 100644 index 000000000..665380e2d --- /dev/null +++ b/src/components/agent/chat/utils/modelThemePolicy.test.ts @@ -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 { + 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"); + }); +}); diff --git a/src/components/agent/chat/utils/modelThemePolicy.ts b/src/components/agent/chat/utils/modelThemePolicy.ts new file mode 100644 index 000000000..e1ca97b09 --- /dev/null +++ b/src/components/agent/chat/utils/modelThemePolicy.ts @@ -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", + }; +} diff --git a/vite.config.ts b/vite.config.ts index 400f9d890..a54d0df07 100644 --- a/vite.config.ts +++ b/vite.config.ts @@ -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/**"], },