Skip to content

Commit 48e9247

Browse files
onderycursoragent
andcommitted
fix(models): pass selected models into verify and keep SUCCESS on later type fail
Two remaining verify-path fixes after infiniflow#17158 (api_base -> base_url rename) and infiniflow#17364/infiniflow#17397 (empty-catalog remote model discovery) landed on main: - useVerifyProvider now accepts a modelInfoRef carrying the models selected in the List-models picker (not registered as form fields) and folds them into the verify payload, so local/compatible providers no longer verify against an empty model_info. - _record_model_verify_failure records FAIL without overwriting a prior SUCCESS for the same llm_name. model_info is unrolled into one factory_llms entry per capability type, so a later type failure must not erase an earlier successful capability result. - Wrap Embedding/Chat model construction in try/except so an init failure is recorded as a per-model FAIL instead of aborting the whole verify. Co-authored-by: Cursor <cursoragent@cursor.com>
1 parent b8bb229 commit 48e9247

3 files changed

Lines changed: 73 additions & 27 deletions

File tree

api/apps/services/provider_api_service.py

Lines changed: 47 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -614,6 +614,17 @@ async def _run_verification(label: str, coro, timeout_seconds: int):
614614
return False, f"\nFail to access {label}.{str(e)}"
615615

616616

617+
def _record_model_verify_failure(model_verify_result: dict, llm_name: str) -> None:
618+
"""Record FAIL without overwriting a prior SUCCESS for the same model.
619+
620+
model_info is unrolled into one factory_llms entry per capability type, so
621+
the same llm_name can be verified multiple times. A later type failure must
622+
not erase an earlier successful capability result.
623+
"""
624+
if model_verify_result.get(llm_name) != ModelVerifyStatusEnum.SUCCESS.value:
625+
model_verify_result[llm_name] = ModelVerifyStatusEnum.FAIL.value
626+
627+
617628
async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url: str = None, region: str = None, model_info: list[dict] = None):
618629
"""
619630
Verify API key for a provider.
@@ -711,27 +722,40 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
711722
if mt_value == LLMType.EMBEDDING.value:
712723
if provider_name not in EmbeddingModel:
713724
msg += f"\nEmbedding model from {provider_name} is not supported yet."
714-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
725+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
715726
continue
716-
mdl = EmbeddingModel[provider_name](api_key_str, llm["llm_name"], base_url=base_url)
717727
label = f"embedding model({llm['llm_name']})"
728+
try:
729+
mdl = EmbeddingModel[provider_name](api_key_str, llm["llm_name"], base_url=base_url)
730+
except Exception as e:
731+
logging.exception("Fail to init %s", label)
732+
msg += f"\nFail to access {label}.{str(e)}"
733+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
734+
continue
718735
ok, result = await _run_verification(label, asyncio.to_thread(mdl.encode, ["Test if the api key is available"]), timeout_seconds)
719736
if not ok:
720737
msg += result
721-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
738+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
722739
continue
723740
if len(result[0]) == 0:
724741
msg += f"\nFail to access {label}."
725-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
742+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
726743
continue
727744
passed = True
728745

729746
elif mt_value == LLMType.CHAT.value:
730747
if provider_name not in ChatModel:
731748
msg += f"\nChat model from {provider_name} is not supported yet."
732-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
749+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
750+
continue
751+
label = f"model({provider_name}/{llm['llm_name']})"
752+
try:
753+
mdl = ChatModel[provider_name](api_key_str, llm["llm_name"], base_url=base_url, **extra)
754+
except Exception as e:
755+
logging.exception("Fail to init %s", label)
756+
msg += f"\nFail to access {label}.{str(e)}"
757+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
733758
continue
734-
mdl = ChatModel[provider_name](api_key_str, llm["llm_name"], base_url=base_url, **extra)
735759

736760
temperature = 1 if llm["llm_name"] in ("kimi-k3", "kimi-k2.7-code") else 0.9
737761

@@ -745,60 +769,59 @@ async def check_streamly():
745769
return True
746770
return False
747771

748-
label = f"model({provider_name}/{llm['llm_name']})"
749772
ok, result = await _run_verification(label, check_streamly(), timeout_seconds)
750773
if not ok:
751774
msg += result
752-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
775+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
753776
continue
754777
if not result:
755778
msg += f"\nFail to access {label}.No valid response received"
756-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
779+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
757780
continue
758781
passed = True
759782

760783
elif mt_value == LLMType.RERANK.value:
761784
if provider_name not in RerankModel:
762785
msg += f"\nRerank model from {provider_name} is not supported yet."
763-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
786+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
764787
continue
765788
mdl = RerankModel[provider_name](api_key_str, llm["llm_name"], base_url=base_url)
766789
label = f"model({provider_name}/{llm['llm_name']})"
767790
ok, result = await _run_verification(label, asyncio.to_thread(mdl.similarity, "What's the weather?", ["Is it sunny today?"]), timeout_seconds)
768791
if not ok:
769792
msg += result
770-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
793+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
771794
continue
772795
arr, tc = result
773796
if len(arr) == 0 or tc == 0:
774797
msg += f"\nFail to access {label}."
775-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
798+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
776799
continue
777800
passed = True
778801

779802
elif mt_value == LLMType.OCR.value:
780803
if provider_name not in OcrModel:
781804
msg += f"\nOCR model from {provider_name} is not supported yet."
782-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
805+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
783806
continue
784807
mdl = OcrModel[provider_name](key=api_key_str, model_name=llm["llm_name"], base_url=base_url)
785808
label = f"model({provider_name}/{llm['llm_name']})"
786809
ok, result = await _run_verification(label, asyncio.to_thread(mdl.check_available), timeout_seconds)
787810
if not ok:
788811
msg += result
789-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
812+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
790813
continue
791814
ok2, reason = result
792815
if not ok2:
793816
msg += f"\nFail to access {label}.{reason or 'Model not available'}"
794-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
817+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
795818
continue
796819
passed = True
797820

798821
elif mt_value == LLMType.TTS.value:
799822
if provider_name not in TTSModel:
800823
msg += f"\nTTS model from {provider_name} is not supported yet."
801-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
824+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
802825
continue
803826
mdl = TTSModel[provider_name](key=api_key_str, model_name=llm["llm_name"], base_url=base_url)
804827

@@ -810,14 +833,14 @@ def drain_tts():
810833
ok, result = await _run_verification(label, asyncio.to_thread(drain_tts), timeout_seconds)
811834
if not ok:
812835
msg += result
813-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
836+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
814837
continue
815838
passed = True
816839

817840
elif mt_value == LLMType.VISION.value:
818841
if provider_name not in CvModel:
819842
msg += f"\nImage to text model from {provider_name} is not supported yet."
820-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
843+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
821844
continue
822845
from rag.utils.base64_image import test_image
823846

@@ -826,31 +849,31 @@ def drain_tts():
826849
ok, result = await _run_verification(label, asyncio.to_thread(mdl.describe, test_image), timeout_seconds)
827850
if not ok:
828851
msg += result
829-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
852+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
830853
continue
831854
m, tc = result
832855
if not tc and m.find("**ERROR**:") >= 0:
833856
msg += f"\nFail to access {label}.{m}"
834-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
857+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
835858
continue
836859
passed = True
837860

838861
elif mt_value == LLMType.ASR.value:
839862
if provider_name not in Seq2txtModel:
840863
msg += f"\nSpeech model from {provider_name} is not supported yet."
841-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
864+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
842865
continue
843866
mdl = Seq2txtModel[provider_name](key=api_key_str, model_name=llm["llm_name"], base_url=base_url)
844867
label = f"model({provider_name}/{llm['llm_name']})"
845868
ok, result = await _run_verification(label, asyncio.to_thread(mdl.check_available), timeout_seconds)
846869
if not ok:
847870
msg += result
848-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
871+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
849872
continue
850873
ok2, reason = result
851874
if not ok2:
852875
msg += f"\nFail to access {label}.{reason or 'Model not available'}"
853-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
876+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
854877
continue
855878
passed = True
856879

@@ -861,7 +884,7 @@ def drain_tts():
861884
any_passed = True
862885
break
863886
else:
864-
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
887+
_record_model_verify_failure(model_verify_result, llm["llm_name"])
865888
if any_passed:
866889
msg = ""
867890

web/src/pages/user-setting/setting-model/instance-card/hooks.tsx

Lines changed: 25 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -305,17 +305,30 @@ type VerifyTransform = (values: Record<string, any>) => {
305305
* `opendataloader_api_key`), it is used to build the verify args;
306306
* otherwise the generic `values.api_key` / `values.base_url` mapping
307307
* is used.
308+
*
309+
* `modelInfoRef` carries the models selected in the List-models picker
310+
* (not registered as form fields). Without it, verify payloads for
311+
* local providers often omit `model_info` and the backend falls back
312+
* to an empty factory catalog.
308313
*/
309314
export function useVerifyProvider(
310315
providerName: string,
311316
formRef: RefObject<DynamicFormRef>,
312317
verifyTransform?: VerifyTransform,
318+
modelInfoRef?: { current: IModelInfo[] },
313319
) {
314320
const { verifyProviderConnection } = useVerifyProviderConnection();
315321

316322
return useCallback(
317323
async (params: any) => {
318324
const values = { ...(formRef.current?.getValues?.() ?? {}), ...params };
325+
const selectedModels =
326+
modelInfoRef?.current?.length && modelInfoRef.current.length > 0
327+
? modelInfoRef.current
328+
: undefined;
329+
if (selectedModels && !values.model_info) {
330+
values.model_info = selectedModels;
331+
}
319332
let verifyArgs: {
320333
api_key: string | object;
321334
base_url?: string;
@@ -327,14 +340,17 @@ export function useVerifyProvider(
327340
verifyArgs = {
328341
api_key: transformed.apiKey,
329342
base_url: transformed.baseUrl,
330-
model_info: transformed.modelInfo ?? values.model_info,
343+
model_info:
344+
transformed.modelInfo?.length
345+
? transformed.modelInfo
346+
: (selectedModels ?? values.model_info),
331347
region: transformed.region,
332348
};
333349
} else {
334350
verifyArgs = {
335351
api_key: values.api_key ?? '',
336352
base_url: values.base_url,
337-
model_info: values.model_info,
353+
model_info: selectedModels ?? values.model_info,
338354
};
339355
}
340356
const ret = await verifyProviderConnection({
@@ -352,7 +368,13 @@ export function useVerifyProvider(
352368
logs: string;
353369
};
354370
},
355-
[providerName, formRef, verifyProviderConnection, verifyTransform],
371+
[
372+
providerName,
373+
formRef,
374+
verifyProviderConnection,
375+
verifyTransform,
376+
modelInfoRef,
377+
],
356378
);
357379
}
358380

web/src/pages/user-setting/setting-model/instance-card/provider-instance-card.tsx

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -158,6 +158,7 @@ const GenericProviderInstanceCard = forwardRef<
158158
providerName,
159159
formRef,
160160
providerConfig.verifyTransform,
161+
modelInfoRef,
161162
);
162163
const handleDelete = useDeleteInstance(
163164
providerName,

0 commit comments

Comments
 (0)