Fix: provider name (#16733)

This commit is contained in:
Lynn
2026-07-09 10:19:10 +08:00
committed by GitHub
parent 575984877f
commit 1430d0e431
65 changed files with 333 additions and 338 deletions

View File

@@ -130,13 +130,7 @@ def _get_model_info(tenant_id: str, default_model: str, model_type: str):
logging.warning(f"Model '{model_name}' is disabled")
return None
return {
"model_provider": provider_name,
"model_instance": instance_name,
"model_name": model_name,
"model_type": model_type,
"enable": True
}
return {"model_provider": provider_name, "model_instance": instance_name, "model_name": model_name, "model_type": model_type, "enable": True}
def _check_model_available(tenant_id: str, provider_name: str, instance_name: str, model_name: str, model_type: str):
@@ -292,14 +286,23 @@ def list_tenant_added_models(tenant_id: str, model_type_filter: str = None):
target_type_records = [record for record in model_records if record.model_type & model_type_filter_bin] if model_type_filter_bin else model_records
factory_rank_mapping = {factory["name"]: -_to_int(factory.get("rank", "500")) for factory in FACTORY_LLM_INFOS}
added_models = [{
"model_type": get_model_type_human(model_record.model_type),
"name": model_record.model_name,
"provider_id": model_record.provider_id,
"provider_name": provider_info_map[model_record.provider_id].provider_name,
"instance_id": model_record.instance_id,
"instance_name": instance_info_map[model_record.instance_id].instance_name
} for model_record in target_type_records]
model_rank_map: dict = {}
for factory in FACTORY_LLM_INFOS:
for llm in factory.get("llm", []):
model_rank_map[(factory["name"], llm["llm_name"])] = _to_int(llm.get("rank", 500))
added_models = [
{
"model_type": get_model_type_human(model_record.model_type),
"name": model_record.model_name,
"provider_id": model_record.provider_id,
"provider_name": provider_info_map[model_record.provider_id].provider_name,
"instance_id": model_record.instance_id,
"instance_name": instance_info_map[model_record.instance_id].instance_name,
"rank": model_rank_map.get((provider_info_map[model_record.provider_id].provider_name, model_record.model_name), 500),
}
for model_record in target_type_records
]
# Add TEI Builtin embedding model if configured
compose_profiles = os.getenv("COMPOSE_PROFILES", "")
@@ -316,10 +319,10 @@ def list_tenant_added_models(tenant_id: str, model_type_filter: str = None):
"provider_name": "Builtin",
"instance_id": "",
"instance_name": "default",
"rank": model_rank_map.get(("Builtin", tei_model), 500),
}
)
added_models.sort(
key=lambda x: (factory_rank_mapping.get(x["provider_name"]), x["provider_name"], x["instance_name"]))
added_models.sort(key=lambda x: (factory_rank_mapping.get(x["provider_name"]), x["provider_name"], x["instance_name"], -x["rank"], x["name"]))
return True, added_models

View File

@@ -20,7 +20,7 @@ import asyncio
from common.constants import LLMType, ActiveStatusEnum, ModelVerifyStatusEnum
from common.settings import FACTORY_LLM_INFOS
from api.db.joint_services.tenant_model_service import get_model_config_from_provider_instance, delete_models_by_instance_ids, delete_instances_by_provider_ids, _decode_api_key_config
from api.db.joint_services.tenant_model_service import get_model_config_from_provider_instance, delete_models_by_instance_ids, delete_instances_by_provider_ids
from api.db.services.tenant_model_provider_service import TenantModelProviderService
from api.db.services.tenant_model_instance_service import TenantModelInstanceService
from api.db.services.tenant_model_service import TenantModelService
@@ -214,7 +214,7 @@ async def list_provider_models(provider_id_or_name: str, api_key: str = None, ba
remote_models = await ModelMeta[provider_name](api_key, model_base_url).get_model_list()
if not static_llms and not remote_models:
return False, f"No models found for provider '{provider_id_or_name}'"
return True, []
# Merge static and remote models, preferring remote_models on name conflicts
merged = {m["name"]: m for m in static_llms}
@@ -257,7 +257,9 @@ def show_provider_model(provider_id_or_name: str, model_name: str):
}
async def update_provider_instance(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, instance_name: str, api_key: str|dict, base_url: str, region: str, model_info: list[dict]=None, verify: bool=True):
async def update_provider_instance(
tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, instance_name: str, api_key: str | dict, base_url: str, region: str, model_info: list[dict] = None, verify: bool = True
):
"""
Update a provider instance.
@@ -420,12 +422,9 @@ async def update_provider_instance(tenant_id: str, provider_id_or_name: str, ins
if verify:
verify_status = model_verify_result.get(llm_name, ModelVerifyStatusEnum.UNKNOWN.value)
extra_fields["verify"] = verify_status
success, _msg = add_model_to_instance(tenant_id, provider_name, effective_instance_name, **{
"model_type": _factory_model_types(llm),
"model_name": llm_name,
"max_tokens": llm["max_tokens"],
"extra": extra_fields
})
success, _msg = add_model_to_instance(
tenant_id, provider_name, effective_instance_name, **{"model_type": _factory_model_types(llm), "model_name": llm_name, "max_tokens": llm["max_tokens"], "extra": extra_fields}
)
if not success:
msg += _msg
@@ -511,16 +510,21 @@ async def create_provider_instance(tenant_id: str, provider_id_or_name: str, ins
factory_llms = factory_info[0]["llm"]
for llm in factory_llms:
llm_name = _factory_llm_name(llm)
success, _msg = add_model_to_instance(tenant_id, provider_name, instance_name, **{
"model_type": _factory_model_types(llm),
"model_name": llm_name,
"max_tokens": llm["max_tokens"],
"extra": {
"is_tools": llm.get("is_tools", False),
"thinking": "thinking" in llm.get("features", []),
"verify": model_verify_result.get(llm_name, ModelVerifyStatusEnum.UNKNOWN.value)
}
})
success, _msg = add_model_to_instance(
tenant_id,
provider_name,
instance_name,
**{
"model_type": _factory_model_types(llm),
"model_name": llm_name,
"max_tokens": llm["max_tokens"],
"extra": {
"is_tools": llm.get("is_tools", False),
"thinking": "thinking" in llm.get("features", []),
"verify": model_verify_result.get(llm_name, ModelVerifyStatusEnum.UNKNOWN.value),
},
},
)
if not success:
msg += _msg
if msg:
@@ -552,12 +556,7 @@ async def create_name_only_provider_instance(tenant_id: str, provider_name: str,
if not provider_obj:
return False, f"Provider '{provider_name}' does not exist"
TenantModelInstanceService.create_instance(
provider_id=provider_obj.id,
instance_name=instance_name,
api_key="",
extra=json.dumps({})
)
TenantModelInstanceService.create_instance(provider_id=provider_obj.id, instance_name=instance_name, api_key="", extra=json.dumps({}))
return True, "success"
@@ -634,10 +633,15 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
return False, f"Provider '{provider_id_or_name}' not found", {}
if model_info:
factory_llms = [{
"model_type": _type,
"llm_name": model.get("model_name", ""),
} for model in model_info if model for _type in model.get("model_type", [])]
factory_llms = [
{
"model_type": _type,
"llm_name": model.get("model_name", ""),
}
for model in model_info
if model
for _type in model.get("model_type", [])
]
if not factory_llms:
return False, f"No valid models found for provider '{provider_id_or_name}'", {}
else:
@@ -820,7 +824,7 @@ def show_provider_instance(tenant_id: str, provider_id_or_name: str, instance_id
"region": extra_fields.get("region", ""),
"base_url": extra_fields.get("base_url", ""),
"api_key": instance_obj.api_key,
"status": instance_obj.status
"status": instance_obj.status,
}
@@ -888,10 +892,17 @@ def list_instance_models(tenant_id: str, provider_id_or_name: str, instance_id_o
if not factory_info:
return False, f"Provider '{provider_id_or_name}' not found"
llms = factory_info[0].get("llm", [])
models = [{"name": llm["llm_name"]} for llm in llms]
models.sort(key=lambda x: x["name"])
models = [{"name": llm["llm_name"], "rank": _to_int(llm.get("rank", 500))} for llm in llms]
models.sort(key=lambda x: (-x["rank"], x["name"]))
return True, models
# Build rank mapping from LLM dictionary for the provider
factory_info = [f for f in FACTORY_LLM_INFOS if f["name"] == provider_obj.provider_name]
model_rank_map = {}
if factory_info:
for llm in factory_info[0].get("llm", []):
model_rank_map[llm["llm_name"]] = _to_int(llm.get("rank", 500))
# Get instance
instance_obj = None
if instance_id_or_name:
@@ -907,14 +918,18 @@ def list_instance_models(tenant_id: str, provider_id_or_name: str, instance_id_o
model_list = []
for model in model_objs:
model_extra = json.loads(model.extra)
model_list.append({
"name": model.model_name,
"model_type": get_model_type_human(model.model_type),
"max_tokens": model_extra.get("max_tokens", 8192) if model.extra else 8192,
"status": model.status,
"verify": model_extra.get("verify", ModelVerifyStatusEnum.UNKNOWN.value),
"features": (["is_tools"] if model_extra.get("is_tools") else []) + (["thinking"] if model_extra.get("thinking") else [])
})
model_list.append(
{
"name": model.model_name,
"model_type": get_model_type_human(model.model_type),
"max_tokens": model_extra.get("max_tokens", 8192) if model.extra else 8192,
"status": model.status,
"verify": model_extra.get("verify", ModelVerifyStatusEnum.UNKNOWN.value),
"features": (["is_tools"] if model_extra.get("is_tools") else []) + (["thinking"] if model_extra.get("thinking") else []),
"rank": model_rank_map.get(model.model_name, 500),
}
)
model_list.sort(key=lambda x: (-x["rank"], x["name"]))
return True, model_list
@@ -984,13 +999,7 @@ def add_model_to_instance(tenant_id: str, provider_id_or_name: str, instance_id_
extra_fields.update({"thinking": "thinking" in target_model[0].get("features", [])})
if extra:
extra_fields.update(extra)
TenantModelService.insert(
model_name=model_name,
provider_id=provider_obj.id,
instance_id=instance_obj.id,
model_type=model_type_bin,
extra=json.dumps(extra_fields)
)
TenantModelService.insert(model_name=model_name, provider_id=provider_obj.id, instance_id=instance_obj.id, model_type=model_type_bin, extra=json.dumps(extra_fields))
return True, "success"

View File

@@ -1210,12 +1210,7 @@ def get_pending_task_count(priority=None):
query = (
Task.select(fn.COUNT(Task.id))
.join(Document, on=(Task.doc_id == Document.id))
.where(
(Task.progress == 0)
& ((Document.run.is_null(True)) | (Document.run != TaskStatus.CANCEL.value))
& (Document.progress >= 0)
& (Document.status == StatusEnum.VALID.value)
)
.where((Task.progress == 0) & ((Document.run.is_null(True)) | (Document.run != TaskStatus.CANCEL.value)) & (Document.progress >= 0) & (Document.status == StatusEnum.VALID.value))
)
if priority is not None:
query = query.where(Task.priority == priority)

View File

@@ -109,8 +109,7 @@ class MemoryService(CommonService):
@classmethod
@DB.connection_context()
def create_memory(cls, tenant_id: str, name: str, memory_type: List[str], embd_id: str, llm_id: str,
tenant_embd_id: str | None = None, tenant_llm_id: str | None = None):
def create_memory(cls, tenant_id: str, name: str, memory_type: List[str], embd_id: str, llm_id: str, tenant_embd_id: str | None = None, tenant_llm_id: str | None = None):
# Deduplicate name within tenant
memory_name = duplicate_name(cls.query, name=name, tenant_id=tenant_id)
if len(memory_name) > MEMORY_NAME_LIMIT: