diff --git a/src/pages/maas-provider/hooks/use-register-route.ts b/src/pages/maas-provider/hooks/use-register-route.ts index 794b15cb..5245e030 100644 --- a/src/pages/maas-provider/hooks/use-register-route.ts +++ b/src/pages/maas-provider/hooks/use-register-route.ts @@ -15,9 +15,8 @@ const useRegisterRoute = () => { routeTargets: [ { weight: 100, - provider_model_name: model.name, - provider_id: record.id, - parentId: record.id + overridden_model_name: model.name, + provider_id: record.id } ] }; diff --git a/src/pages/model-routes/components/route-targets.tsx b/src/pages/model-routes/components/route-targets.tsx index 4e4cec32..361c3b13 100644 --- a/src/pages/model-routes/components/route-targets.tsx +++ b/src/pages/model-routes/components/route-targets.tsx @@ -78,7 +78,7 @@ const RouteItem: React.FC = ({ {data.model_id ? modelList?.find((m) => m.value === data.model_id)?.label - : data.provider_model_name} + : data.overridden_model_name} ); diff --git a/src/pages/model-routes/config/types.ts b/src/pages/model-routes/config/types.ts index 8266cfc5..902f9ffc 100644 --- a/src/pages/model-routes/config/types.ts +++ b/src/pages/model-routes/config/types.ts @@ -1,24 +1,21 @@ +export interface RouteTargetFormItem { + id?: number; + overridden_model_name?: string; + weight?: number | null; + model_id?: number; + provider_id?: number; + fallback_status_codes?: string[]; + parentId?: string | number; +} + export interface FormData { name: string; description: string; categories: any[]; meta: Record; generic_proxy: boolean; - fallback_target: { - provider_model_name?: string; - model_id?: number; - provider_id?: number; - lora_module_name?: string; - fallback_status_codes?: string[]; - }; - targets: { - provider_model_name?: string; - weight?: number | null; - model_id?: number; - provider_id?: number; - lora_module_name?: string; - fallback_status_codes?: string[]; - }[]; + fallback_target: RouteTargetFormItem | null; + targets: RouteTargetFormItem[]; } export interface RouteItem { @@ -36,7 +33,7 @@ export interface RouteItem { access_policy: string; } -export interface RouteTarget { +export interface RouteTarget extends RouteTargetFormItem { id: number; created_at: string; updated_at: string; @@ -46,8 +43,7 @@ export interface RouteTarget { route_name: string; route_id: number; provider_id: number; - provider_model_name: string; + overridden_model_name: string; fallback_status_codes: string[]; - lora_module_name?: string; state: string; } diff --git a/src/pages/model-routes/forms/index.tsx b/src/pages/model-routes/forms/index.tsx index 23696f4c..e2902cab 100644 --- a/src/pages/model-routes/forms/index.tsx +++ b/src/pages/model-routes/forms/index.tsx @@ -14,23 +14,56 @@ import { Form } from 'antd'; import _ from 'lodash'; import { forwardRef, useEffect, useImperativeHandle, useRef } from 'react'; import FormContext from '../config/form-context'; -import { FormData, RouteItem as ListItem } from '../config/types'; +import { + FormData, + RouteItem as ListItem, + RouteTargetFormItem +} from '../config/types'; import useEditTargets from '../hooks/use-edit-targets'; import Basic from './basic'; import Targets from './targets'; +const isSameTarget = ( + left?: RouteTargetFormItem | null, + right?: RouteTargetFormItem | null +) => { + if (!left || !right) { + return false; + } + + if (left.model_id != null || right.model_id != null) { + return left.model_id != null && left.model_id === right.model_id; + } + + return ( + left.provider_id != null && + left.provider_id === right.provider_id && + left.overridden_model_name === right.overridden_model_name + ); +}; + +const normalizeTarget = ( + target: RouteTargetFormItem | null | undefined +): RouteTargetFormItem => { + const normalizedTarget = { + id: target?.id, + weight: target?.weight ?? 0, + model_id: target?.model_id, + provider_id: target?.provider_id, + overridden_model_name: target?.overridden_model_name, + fallback_status_codes: target?.fallback_status_codes + }; + + return _.omitBy(normalizedTarget, _.isUndefined) as RouteTargetFormItem; +}; + interface ProviderFormProps { ref?: any; open: boolean; action: PageActionType; realAction?: string; currentData?: ListItem & { - routeTargets?: { - weight?: number; - model_id?: number; - provider_id?: number; - provider_model_name?: string; - }[]; + routeTargets?: RouteTargetFormItem[]; }; // Used when action is EDIT onFinish: (values: FormData) => Promise; onFallbackChange?: (changed: boolean) => void; @@ -92,25 +125,12 @@ const AccessForm: React.FC = forwardRef((props, ref) => { const fallbackTarget = values.fallback_target; if (fallbackTarget) { - const exsitinged = targetList.find((ep) => { - if (fallbackTarget!.model_id) { - return ( - ep.model_id === fallbackTarget!.model_id && - ep.lora_module_name === fallbackTarget!.lora_module_name - ); - } - return ( - ep.provider_id === fallbackTarget!.provider_id && - ep.provider_model_name === fallbackTarget!.provider_model_name - ); - }); - if (exsitinged) { + const existingTarget = targetList.find((target) => + isSameTarget(target, fallbackTarget) + ); + if (existingTarget) { targetList = targetList.map((ep) => { - if ( - (ep.model_id === fallbackTarget.model_id && - ep.lora_module_name === fallbackTarget.lora_module_name) || - ep.provider_model_name === fallbackTarget.provider_model_name - ) { + if (isSameTarget(ep, fallbackTarget)) { return { ...ep, fallback_status_codes: ['4xx', '5xx'] @@ -120,7 +140,7 @@ const AccessForm: React.FC = forwardRef((props, ref) => { }); } - if (!exsitinged) { + if (!existingTarget) { targetList.push({ ...fallbackTarget, weight: 0, @@ -129,7 +149,7 @@ const AccessForm: React.FC = forwardRef((props, ref) => { } } - return targetList; + return targetList.map((target) => normalizeTarget(target)); }; const handleOnFinish = (values: FormData) => { @@ -148,27 +168,14 @@ const AccessForm: React.FC = forwardRef((props, ref) => { return; } - const initDataList = async ( - targets: { - weight?: number; - model_id?: number; - provider_id?: number; - provider_model_name?: string; - lora_module_name?: string; - }[] - ) => { + const initDataList = async (targets: RouteTargetFormItem[]) => { // init targets form list targetsRef.current?.initDataList( targets?.map((ep) => ({ weight: ep.weight, value: ep.model_id - ? [ - 'deployments', - ep.lora_module_name - ? `${ep.model_id}_lora_${ep.lora_module_name}` - : ep.model_id - ] - : [ep.provider_id, ep.provider_model_name] + ? ['deployments', ep.model_id] + : [ep.provider_id, ep.overridden_model_name] })) || [] ); }; @@ -190,13 +197,8 @@ const AccessForm: React.FC = forwardRef((props, ref) => { if (fallbackTarget) { targetsRef.current?.initFallbackValues({ value: fallbackTarget.model_id - ? [ - 'deployments', - fallbackTarget.lora_module_name - ? `${fallbackTarget.model_id}_lora_${fallbackTarget.lora_module_name}` - : fallbackTarget.model_id - ] - : [fallbackTarget.provider_id, fallbackTarget.provider_model_name] + ? ['deployments', fallbackTarget.model_id] + : [fallbackTarget.provider_id, fallbackTarget.overridden_model_name] }); } }; diff --git a/src/pages/model-routes/hooks/use-target-source-models.tsx b/src/pages/model-routes/hooks/use-target-source-models.tsx index 034cab7f..7f083975 100644 --- a/src/pages/model-routes/hooks/use-target-source-models.tsx +++ b/src/pages/model-routes/hooks/use-target-source-models.tsx @@ -90,7 +90,7 @@ const useTargetSourceModels = () => { label: model.name, value: model.name, data: { - provider_model_name: model.name, + overridden_model_name: model.name, provider_id: provider.id, parentId: provider.id },