import React, { useEffect, useMemo, useRef, useState } from "react"; import { AutoComplete, Button, Card, Col, Divider, Drawer, Form, Input, InputNumber, Popconfirm, Row, Select, Space, Switch, Table, Tabs, Tag, Tooltip, Typography, App } from 'antd'; import PageContainer from "@/components/shared/PageContainer"; import { DeleteOutlined, EditOutlined, PlusOutlined, SafetyCertificateOutlined, SaveOutlined, SearchOutlined, SyncOutlined, WifiOutlined, } from "@ant-design/icons"; import { useDict } from "../../hooks/useDict"; import { AiModelDTO, AiLocalProfileVO, AiModelVO, deleteAiModelByType, getAiModelPage, getRemoteModelList, saveAiModel, testLlmModelConnectivity, testLocalModelConnectivity, updateAiModel, } from "../../api/business/aimodel"; import AppPagination from "../../components/shared/AppPagination"; const { Option } = Select; const { Title } = Typography; type ModelType = "ASR" | "LLM"; const PROVIDER_BASE_URL_MAP: Record = { openai: "https://api.openai.com", deepseek: "https://api.deepseek.com", aliyun: "https://dashscope.aliyuncs.com/compatible-mode", qwen: "https://dashscope.aliyuncs.com/compatible-mode", dashscope: "https://dashscope.aliyuncs.com/compatible-mode", moonshot: "https://api.moonshot.cn", kimi: "https://api.moonshot.cn", groq: "https://api.groq.com/openai", }; const DEFAULT_LLM_TEST_MESSAGE = "请回复:LLM 连通性测试成功。"; const AiModels: React.FC = () => { const { message } = App.useApp(); const [form] = Form.useForm(); const { items: providers } = useDict("biz_ai_provider"); const [activeType, setActiveType] = useState("ASR"); const [loading, setLoading] = useState(false); const [data, setData] = useState([]); const [total, setTotal] = useState(0); const [current, setCurrent] = useState(1); const [size, setSize] = useState(10); const [searchName, setSearchName] = useState(""); const [drawerVisible, setDrawerVisible] = useState(false); const [editingId, setEditingId] = useState(null); const [submitLoading, setSubmitLoading] = useState(false); const [fetchLoading, setFetchLoading] = useState(false); const [connectivityLoading, setConnectivityLoading] = useState(false); const [remoteModels, setRemoteModels] = useState([]); const [speakerModels, setSpeakerModels] = useState([]); const modelNameAutoFilledRef = useRef(false); const localProfileLoadedRef = useRef(false); const provider = Form.useWatch("provider", form); const isDefaultChecked = Form.useWatch("isDefaultChecked", form); const isLocalProvider = String(provider || "").toLowerCase() === "custom"; const isTencentProvider = String(provider || "").toLowerCase() === "tencent"; const isPlatformAdmin = useMemo(() => { const profileStr = sessionStorage.getItem("userProfile"); if (!profileStr) { return false; } const profile = JSON.parse(profileStr); return profile.isPlatformAdmin === true; }, []); useEffect(() => { void fetchData(); }, [current, size, searchName, activeType]); useEffect(() => { if (!drawerVisible || !provider) { return; } const providerItem = providers.find((item) => item.itemValue === provider); const providerLabel = providerItem?.itemLabel || provider; const currentDisplayName = form.getFieldValue("modelName"); if (!editingId && (!currentDisplayName || modelNameAutoFilledRef.current)) { form.setFieldValue("modelName", providerLabel); modelNameAutoFilledRef.current = true; } const baseUrl = form.getFieldValue("baseUrl"); const providerKey = String(provider).toLowerCase(); const defaultBaseUrl = PROVIDER_BASE_URL_MAP[providerKey]; if (!baseUrl && defaultBaseUrl) { form.setFieldValue("baseUrl", defaultBaseUrl); } }, [provider, drawerVisible, editingId, providers, form]); const fetchData = async () => { setLoading(true); try { const res = await getAiModelPage({ current, size, name: searchName || undefined, type: activeType, }); const pageData = (res as any)?.data?.data ?? (res as any); setData(pageData?.records || []); setTotal(pageData?.total || 0); } finally { setLoading(false); } }; const openDrawer = (record?: AiModelVO) => { setRemoteModels([]); setSpeakerModels([]); modelNameAutoFilledRef.current = false; localProfileLoadedRef.current = false; if (record) { setEditingId(record.id); const speakerModel = record.mediaConfig?.speakerModel; const svThreshold = record.mediaConfig?.svThreshold; const tencentAppId = record.mediaConfig?.tencentAppId; const tencentSecretId = record.mediaConfig?.tencentSecretId; const tencentSecretKey = record.mediaConfig?.tencentSecretKey; form.setFieldsValue({ ...record, modelType: record.modelType, speakerModel, svThreshold, tencentAppId, tencentSecretId, tencentSecretKey, isDefaultChecked: record.isDefault === 1, statusChecked: record.status === 1, }); if (record.modelCode) { setRemoteModels([record.modelCode]); } if (speakerModel) { setSpeakerModels([String(speakerModel)]); } } else { setEditingId(null); form.resetFields(); form.setFieldsValue({ modelType: activeType, isDefaultChecked: false, statusChecked: true, sortOrder: 0, temperature: 0.2, topP: 0.9, apiPath: "/v1/chat/completions", svThreshold: 0.45, }); } setDrawerVisible(true); }; useEffect(() => { if (!drawerVisible || !isLocalProvider || !editingId || localProfileLoadedRef.current) { return; } const values = form.getFieldsValue(["baseUrl", "apiKey"]); if (!values.baseUrl || !values.apiKey) { return; } localProfileLoadedRef.current = true; void handleTestConnectivity(); }, [drawerVisible, isLocalProvider, editingId, form]); const handleFetchRemote = async () => { if (isLocalProvider) { await handleTestConnectivity(); return; } const values = form.getFieldsValue(["provider", "baseUrl", "apiKey"]); if (!values.provider || !values.baseUrl) { message.warning("请先填写提供商和 Base URL"); return; } setFetchLoading(true); try { const res = await getRemoteModelList(values); const rawModels = (res as any)?.data?.data ?? (Array.isArray(res) ? res : []); const models = Array.isArray(rawModels) ? rawModels : []; setRemoteModels(models); message.success(`获取到 ${models.length} 个模型`); } finally { setFetchLoading(false); } }; const resolveLocalWsUrl = (baseUrl: string, wsEndpoint?: string) => { if (!wsEndpoint) { return undefined; } try { const base = new URL(baseUrl); const endpoint = wsEndpoint.startsWith("/") ? wsEndpoint : `/${wsEndpoint}`; const protocol = base.protocol === "https:" ? "wss:" : "ws:"; return `${protocol}//${base.host}${endpoint}`; } catch { return undefined; } }; const applyLocalProfile = (profile: AiLocalProfileVO, baseUrl: string) => { const nextRemoteModels = Array.isArray(profile.asrModels) ? profile.asrModels : []; const nextSpeakerModels = Array.isArray(profile.speakerModels) ? profile.speakerModels : []; setRemoteModels(nextRemoteModels); setSpeakerModels(nextSpeakerModels); const nextValues: Record = {}; if (profile.activeAsrModel) { nextValues.modelCode = profile.activeAsrModel; } if (profile.activeSpeakerModel) { nextValues.speakerModel = profile.activeSpeakerModel; } if (profile.svThreshold !== undefined) { nextValues.svThreshold = profile.svThreshold; } const wsUrl = resolveLocalWsUrl(baseUrl, profile.wsEndpoint); if (wsUrl) { nextValues.wsUrl = wsUrl; } form.setFieldsValue(nextValues); }; const handleSubmit = async () => { const values = await form.validateFields(); if (values.isDefaultChecked && !values.statusChecked) { message.warning("默认模型必须保持启用状态"); return; } const payload: AiModelDTO = { id: editingId ?? undefined, modelType: values.modelType, modelName: values.modelName, provider: values.provider, baseUrl: values.baseUrl, apiPath: values.apiPath, apiKey: values.apiKey, modelCode: values.modelCode, wsUrl: values.wsUrl, mediaConfig: activeType === "ASR" && isLocalProvider ? { speakerModel: values.speakerModel, svThreshold: values.svThreshold, } : activeType === "ASR" && isTencentProvider ? { tencentAppId: values.tencentAppId, tencentSecretId: values.tencentSecretId, tencentSecretKey: values.tencentSecretKey, } : undefined, temperature: values.temperature, topP: values.topP, isDefault: values.isDefaultChecked ? 1 : 0, status: values.statusChecked ? 1 : 0, sortOrder: values.sortOrder ?? 0, remark: values.remark, }; setSubmitLoading(true); try { if (editingId) { await updateAiModel(payload); message.success("更新成功"); } else { await saveAiModel(payload); message.success("新增成功"); } setDrawerVisible(false); void fetchData(); } finally { setSubmitLoading(false); } }; const handleTestConnectivity = async () => { if (activeType === "LLM") { const values = await form.validateFields(["provider", "baseUrl", "apiPath", "modelCode"]); const extraValues = form.getFieldsValue(["apiKey", "temperature", "topP"]); setConnectivityLoading(true); try { await testLlmModelConnectivity({ provider: values.provider, baseUrl: values.baseUrl, apiPath: values.apiPath, apiKey: extraValues.apiKey, modelCode: values.modelCode, temperature: extraValues.temperature, topP: extraValues.topP, testMessage: DEFAULT_LLM_TEST_MESSAGE, }); message.success("LLM 连通性测试成功"); } finally { setConnectivityLoading(false); } return; } const values = await form.validateFields(["provider", "baseUrl"]); if (String(values.provider || "").toLowerCase() !== "custom") { message.warning("仅本地模型支持连通性测试"); return; } const { apiKey } = form.getFieldsValue(["apiKey"]); if (!apiKey) { message.warning("请先填写 API Key 后再测试连接"); return; } setConnectivityLoading(true); try { const res = await testLocalModelConnectivity({ baseUrl: values.baseUrl, apiKey, }); const profile = (res as any)?.data?.data ?? (res as any)?.data ?? (res as any); applyLocalProfile(profile as AiLocalProfileVO, values.baseUrl); message.success("本地模型连通性测试成功"); } finally { setConnectivityLoading(false); } }; const handleDelete = async (record: AiModelVO) => { await deleteAiModelByType(record.id, record.modelType); message.success("删除成功"); void fetchData(); }; const columns = [ { title: "模型名称", dataIndex: "modelName", key: "modelName", render: (text: string, record: AiModelVO) => ( {text} {record.isDefault === 1 && 默认} {record.tenantId === 0 && ( )} ), }, { title: "提供商", dataIndex: "provider", key: "provider", render: (value: string) => { const item = providers.find((providerItem) => providerItem.itemValue === value); return item ? {item.itemLabel} : value; }, }, { title: "模型名称(code)", dataIndex: "modelCode", key: "modelCode" }, { title: "排序", dataIndex: "sortOrder", key: "sortOrder", render: (value: number | undefined) => value ?? 0, }, { title: "状态", dataIndex: "status", key: "status", render: (status: number) => status === 1 ? 启用 : 禁用, }, { title: "操作", key: "action", render: (_: unknown, record: AiModelVO) => { const canEdit = record.tenantId !== 0 || isPlatformAdmin; return ( {canEdit && ( )} {canEdit && ( handleDelete(record)}> )} ); }, }, ]; return ( } onClick={() => openDrawer()}> 新增模型 } toolbar={ } allowClear onPressEnter={(event) => setSearchName((event.target as HTMLInputElement).value)} style={{ width: 220 }} /> } > { setActiveType(key as ModelType); setCurrent(1); }} items={[ { key: "ASR", label: "ASR 模型" }, { key: "LLM", label: "LLM 模型" }, ]} style={{ marginBottom: 16 }} />
{ setCurrent(page); setSize(pageSize); }} /> setDrawerVisible(false)} title={{editingId ? "编辑模型" : "新增模型"}} forceRender extra={ } >
{activeType === "ASR" ? "语音识别 (ASR)" : "总结模型 (LLM)"}
{ modelNameAutoFilledRef.current = false; }} /> {!isTencentProvider && ( <> )} {(activeType === "LLM" || isLocalProvider) && ( )} 模型参数 { if (isLocalProvider && remoteModels.length === 0) { void handleFetchRemote(); } }} options={remoteModels.map((model) => ({ value: model }))} filterOption={(inputValue, option) => isLocalProvider || String(option?.value || "").toLowerCase().includes(inputValue.toLowerCase()) } > {!isTencentProvider && ( )} {activeType === "ASR" && isLocalProvider && ( */} {/* */} {/* ({ label: model, value: model }))}*/} {/* />*/} {/* */} {/**/} )} {activeType === "ASR" && isTencentProvider && ( )} {activeType === "LLM" && ( <> )} { if (checked) { form.setFieldValue("statusChecked", true); } }} /> ); }; export default AiModels;