Feat: v0.27.0 model provider (#16604)

This commit is contained in:
Lynn
2026-07-08 09:47:29 +08:00
committed by GitHub
parent cb93883f3f
commit 0ae5961e1c
94 changed files with 7539 additions and 3044 deletions

View File

@@ -23,6 +23,7 @@ from common.exceptions import ArgumentException, NotFoundException
from api.apps import login_required, current_user
from api.utils.api_utils import validate_request, get_request_json, get_error_argument_result, get_json_result
from api.apps.services import memory_api_service
from api.db.joint_services.tenant_model_service import ensure_tenant_model_ids_for_params
from api.utils.pagination_utils import validate_rest_api_page_size
@@ -35,7 +36,15 @@ async def create_memory():
req = await get_request_json()
t_parsed = time.perf_counter() if timing_enabled else None
try:
memory_info = {"name": req["name"], "memory_type": req["memory_type"], "embd_id": req["embd_id"], "llm_id": req["llm_id"]}
# Resolve tenant_model IDs from model names
tenant_id = current_user.id
memory_info = {
"name": req["name"],
"memory_type": req["memory_type"],
"embd_id": req["embd_id"],
"llm_id": req["llm_id"]
}
ensure_tenant_model_ids_for_params(tenant_id, memory_info)
success, res = await memory_api_service.create_memory(memory_info)
if timing_enabled:
logging.info(
@@ -79,26 +88,12 @@ async def create_memory():
@login_required
async def update_memory(memory_id):
req = await get_request_json()
new_settings = {
k: req[k]
for k in [
"name",
"permissions",
"llm_id",
"embd_id",
"memory_type",
"memory_size",
"forgetting_policy",
"temperature",
"avatar",
"description",
"system_prompt",
"user_prompt",
"tenant_llm_id",
"tenant_embd_id",
]
if k in req
}
# Resolve tenant_model IDs from model names when name is provided but id is not
ensure_tenant_model_ids_for_params(current_user.id, req)
new_settings = {k: req[k] for k in [
"name", "permissions", "llm_id", "embd_id", "memory_type", "memory_size", "forgetting_policy", "temperature",
"avatar", "description", "system_prompt", "user_prompt", "tenant_llm_id", "tenant_embd_id"
] if k in req}
try:
success, res = await memory_api_service.update_memory(memory_id, new_settings)
if success:

View File

@@ -17,7 +17,7 @@ import logging
from quart import request
from api.apps import login_required
from api.apps import login_required, current_user
from api.utils.api_utils import (
add_tenant_id_to_kwargs,
get_error_argument_result,
@@ -316,6 +316,9 @@ async def create_provider_instance(tenant_id: str = None, provider_id_or_name: s
required:
- instance_name
- api_key
- base_url
- region
- model_info
properties:
instance_name:
type: string
@@ -336,10 +339,24 @@ async def create_provider_instance(tenant_id: str = None, provider_id_or_name: s
type: object
"""
data = await request.get_json()
if not provider_id_or_name:
return get_error_argument_result(message="provider_id_or_name is required")
if not data or "instance_name" not in data:
return get_error_argument_result(message="instance_name is required")
instance_name = data["instance_name"]
# data only contains instance_name — no other fields needed
if set(data.keys()) == {"instance_name"}:
try:
success, msg = await provider_api_service.create_name_only_provider_instance(tenant_id, provider_id_or_name, instance_name)
if success:
return get_result(message=msg)
else:
return get_error_data_result(message=msg)
except Exception as e:
logging.exception(e)
return get_error_data_result(message="Internal server error")
api_key = data.get("api_key", "")
base_url = data.get("base_url", "")
region = data.get("region", "")
@@ -397,7 +414,10 @@ async def verify_provider_api_key(provider_id_or_name: str = None):
description: Region.
model_info:
type: object
description: Model info.
description: Model info. optional
instance_id:
type: string
description: Instance ID. optional
responses:
200:
description: Instance created successfully.
@@ -405,6 +425,8 @@ async def verify_provider_api_key(provider_id_or_name: str = None):
type: object
"""
data = await request.get_json()
if not provider_id_or_name:
return get_error_argument_result(message="provider_id_or_name is required")
if not data or ("api_key" not in data and provider_id_or_name != "VLLM"):
return get_error_argument_result(message="api_key is required")
@@ -414,8 +436,16 @@ async def verify_provider_api_key(provider_id_or_name: str = None):
model_info = data.get("model_info", [])
try:
success, msg = await provider_api_service.verify_api_key(provider_id_or_name, api_key, base_url, region, model_info)
success, msg, model_verify_result = await provider_api_service.verify_api_key(provider_id_or_name, api_key, base_url, region, model_info)
if success:
if data.get("instance_id"):
# if instance_id is provided, update the model verify result
instance_id = data["instance_id"]
try:
for model, verify_result in model_verify_result.items():
provider_api_service.update_model(current_user.id, provider_id_or_name, instance_id, model, {"verify": verify_result})
except Exception as e:
logging.exception(e)
return get_result(message=msg)
else:
return get_error_data_result(message=msg)
@@ -512,6 +542,117 @@ def show_provider_instance(tenant_id: str = None, provider_id_or_name: str = Non
return get_error_data_result(message="Internal server error")
@manager.route("/providers/<provider_id_or_name>/instances/<instance_id_or_name>", methods=["PUT"]) # noqa: F821
@login_required
@add_tenant_id_to_kwargs
async def update_provider_instance(tenant_id: str = None, provider_id_or_name: str = None, instance_id_or_name: str = None):
"""
Update a provider instance.
---
tags:
- Providers
security:
- ApiKeyAuth: []
parameters:
- in: path
name: provider_id_or_name
type: string
required: true
description: Provider ID or name.
- in: path
name: instance_id_or_name
type: string
required: true
description: Instance ID or name.
- in: header
name: Authorization
type: string
required: true
description: Bearer token for authentication.
- in: body
name: body
description: Instance update parameters.
required: true
schema:
type: object
required:
- instance_name
- api_key
- base_url
- region
- model_info
properties:
instance_name:
type: string
description: Instance name.
api_key:
type: string
description: API key.
base_url:
type: string
description: Base URL.
region:
type: string
description: Region.
model_info:
type: array
description: List of models to configure for this instance.
items:
type: object
properties:
model_type:
type: array
description: Model types.
model_name:
type: string
description: Model name.
max_tokens:
type: integer
description: Max tokens.
extra:
type: object
description: Extra model info (e.g. is_tools).
verify:
type: boolean
description: Verify api_key and base_url, default true
responses:
200:
description: Instance updated successfully.
schema:
type: object
"""
data = await request.get_json()
if not provider_id_or_name:
return get_error_argument_result(message="provider_id_or_name is required")
if not instance_id_or_name:
return get_error_argument_result(message="instance_id_or_name is required")
if not data:
return get_error_argument_result(message="Request body is required")
required_keys = ["instance_name", "api_key", "base_url", "model_info"]
missing = [k for k in required_keys if k not in data]
if missing:
return get_error_argument_result(message=f"Missing required fields: {', '.join(missing)}")
instance_name = data["instance_name"]
api_key = data["api_key"]
base_url = data["base_url"]
region = data.get("region", "default")
model_info = data["model_info"]
verify = data.get("verify", True)
try:
success, msg = await provider_api_service.update_provider_instance(
tenant_id, provider_id_or_name, instance_id_or_name, instance_name, api_key, base_url, region, model_info, verify
)
if success:
return get_result(message=msg)
else:
return get_error_data_result(message=msg)
except Exception as e:
logging.exception(e)
return get_error_data_result(message="Internal server error")
@manager.route("/providers/<provider_id_or_name>/instances", methods=["DELETE"]) # noqa: F821
@login_required
@add_tenant_id_to_kwargs
@@ -761,12 +902,72 @@ async def add_model_to_instance(tenant_id: str, provider_id_or_name: str, instan
return get_error_data_result(message="Internal server error")
@manager.route("/providers/<provider_id_or_name>/instances/<instance_id_or_name>/models", methods=["DELETE"]) # noqa: F821
@login_required
@add_tenant_id_to_kwargs
async def delete_models_from_instance(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str):
"""
Delete models from an instance.
---
tags:
- Providers
security:
- ApiKeyAuth: []
parameters:
- in: path
name: provider_id_or_name
type: string
required: true
description: Provider ID or name.
- in: path
name: instance_id_or_name
type: string
required: true
description: Instance ID or name.
- in: header
name: Authorization
type: string
required: true
description: Bearer token for authentication.
- in: body
name: body
description: Model details.
required: true
schema:
type: object
required:
- model_name
- model_type
properties:
model_name:
type: list of string
description: Model name.
responses:
200:
description: Model deleted successfully.
"""
data = await request.get_json()
if not data or "model_name" not in data:
return get_error_argument_result(message="model_name is required")
model_name = data["model_name"]
try:
success, result = await provider_api_service.delete_models_from_instance(tenant_id, provider_id_or_name, instance_id_or_name, model_name)
if success:
return get_result(message=result)
else:
return get_error_data_result(message=result)
except Exception as e:
logging.exception(e)
return get_error_data_result(message="Internal server error")
@manager.route("/providers/<provider_id_or_name>/instances/<instance_id_or_name>/models/<path:model_name>", methods=["PATCH"]) # noqa: F821
@login_required
@add_tenant_id_to_kwargs
async def enable_or_disable_model(tenant_id: str = None, provider_id_or_name: str = None, instance_id_or_name: str = None, model_name: str = None):
async def alter_model(tenant_id: str = None, provider_id_or_name: str = None, instance_id_or_name: str = None, model_name: str = None):
"""
Enable or disable a model.
Enable or disable a model, or update max_tokens
---
tags:
- Providers
@@ -813,15 +1014,17 @@ async def enable_or_disable_model(tenant_id: str = None, provider_id_or_name: st
type: object
"""
data = await request.get_json()
if not data or "status" not in data:
return get_error_argument_result(message="status is required")
if not data or ("status" not in data and "max_tokens" not in data):
return get_error_argument_result(message="status or max_tokens required.")
status = data["status"]
if status not in ("active", "inactive"):
update_dict = {k: data[k] for k in ["status", "max_tokens", "model_type", "extra"] if k in data}
if update_dict.get("status") and update_dict["status"] not in ("active", "inactive"):
return get_error_argument_result(message="status must be 'active' or 'inactive'")
try:
success, msg = provider_api_service.update_model_status(tenant_id, provider_id_or_name, instance_id_or_name, model_name, status)
success, msg = provider_api_service.update_model(
tenant_id, provider_id_or_name, instance_id_or_name, model_name, update_dict
)
if success:
return get_result(message=msg)
else:

View File

@@ -79,8 +79,8 @@ async def create_memory(memory_info: dict):
"memory_type": list[str],
"embd_id": str,
"llm_id": str,
"tenant_embd_id": str,
"tenant_llm_id": str
"tenant_embd_id": str | None,
"tenant_llm_id": str | None
}
"""
# check name length
@@ -98,7 +98,15 @@ async def create_memory(memory_info: dict):
if invalid_type:
raise ArgumentException(f"Memory type '{invalid_type}' is not supported.")
memory_type = list(memory_type)
success, res = MemoryService.create_memory(tenant_id=current_user.id, name=memory_name, memory_type=memory_type, embd_id=memory_info["embd_id"], llm_id=memory_info["llm_id"])
success, res = MemoryService.create_memory(
tenant_id=current_user.id,
name=memory_name,
memory_type=memory_type,
embd_id=memory_info["embd_id"],
llm_id=memory_info["llm_id"],
tenant_embd_id=memory_info.get("tenant_embd_id"),
tenant_llm_id=memory_info.get("tenant_llm_id"),
)
if success:
return True, format_ret_data_from_memory(res)
else:

View File

@@ -21,6 +21,7 @@ from api.db.services.tenant_model_instance_service import TenantModelInstanceSer
from api.db.services.tenant_model_provider_service import TenantModelProviderService
from api.db.services.tenant_model_service import TenantModelService
from api.db.services.user_service import TenantService
from api.utils.model_utils import get_model_type_human, calculate_model_type
from common.constants import ActiveStatusEnum, LLMType
from common.settings import FACTORY_LLM_INFOS
@@ -120,37 +121,13 @@ def _get_model_info(tenant_id: str, default_model: str, model_type: str):
logging.warning(f"Instance '{instance_name}' not found for provider '{provider_name}'")
return None
# Check if model is enabled (no TenantModel record or status != inactive means enabled)
model_entity = TenantModelService.get_by_provider_id_and_instance_id_and_model_type_and_model_name(provider_obj.id, instance_obj.id, model_type, model_name)
enable = model_entity is None or model_entity.status == ActiveStatusEnum.ACTIVE.value
if not enable:
model_record = TenantModelService.get_by_provider_id_and_instance_id_and_model_name(provider_obj.id, instance_obj.id, model_name)
if not model_record:
logging.warning(f"Model '{model_name}' not found for provider '{provider_name}' and instance '{instance_name}'")
return None
if model_entity:
return {
"model_provider": provider_name,
"model_instance": instance_name,
"model_name": model_name,
"model_type": model_type,
"enable": enable,
}
# Check if model is in the LLM factory info
factory_info = [f for f in (FACTORY_LLM_INFOS or []) if f["name"] == provider_name]
if not factory_info:
logging.warning(f"Provider '{provider_name}' not found in factory info")
return None
llms = factory_info[0].get("llm", [])
target_llm = [llm for llm in llms if llm["llm_name"] == model_name]
if not target_llm:
logging.warning(f"Model '{model_name}' not found for provider '{provider_name}'")
return None
# Check if the model_type matches
if model_type not in _factory_model_types(target_llm[0]):
logging.warning(f"Model '{model_name}' isn't a {model_type} model")
if not model_record.status == ActiveStatusEnum.ACTIVE.value:
logging.warning(f"Model '{model_name}' is disabled")
return None
return {
@@ -158,7 +135,7 @@ def _get_model_info(tenant_id: str, default_model: str, model_type: str):
"model_instance": instance_name,
"model_name": model_name,
"model_type": model_type,
"enable": enable,
"enable": True
}
@@ -200,21 +177,11 @@ def _check_model_available(tenant_id: str, provider_name: str, instance_name: st
return False, f"Provider '{provider_name}' not found in factory info"
model_type = MODEL_TAG_TO_TYPE.get(model_type, model_type)
# Check if model is disabled
model_entity = TenantModelService.get_by_provider_id_and_instance_id_and_model_type_and_model_name(provider_obj.id, instance_obj.id, model_type, model_name)
if model_entity:
if model_entity.status != ActiveStatusEnum.ACTIVE.value:
return False, f"Model '{model_name}' isn't available"
return True, None
llms = factory_info[0].get("llm", [])
target_llm = [llm for llm in llms if llm["llm_name"] == model_name]
if not target_llm and not model_entity:
return False, f"Model '{model_name}' not found for provider '{provider_name}'"
if target_llm:
if model_type not in _factory_model_types(target_llm[0]):
return False, f"Model '{model_name}' isn't a {model_type} model"
model_record = TenantModelService.get_by_provider_id_and_instance_id_and_model_name(provider_obj.id, instance_obj.id, model_name)
if model_record.status != ActiveStatusEnum.ACTIVE.value:
return False, f"Model '{model_name}' isn't available"
if model_type not in get_model_type_human(model_record.model_type):
return False, f"Model '{model_name}' isn't a {model_type} model"
return True, None
@@ -301,8 +268,7 @@ def list_tenant_added_models(tenant_id: str, model_type_filter: str = None):
ensure_paddleocr_from_env(tenant_id)
ensure_opendataloader_from_env(tenant_id)
if model_type_filter:
model_type_filter = model_type_filter.lower()
model_type_filter_bin = calculate_model_type(model_type_filter.lower()) if model_type_filter else None
providers = TenantModelProviderService.get_by_tenant_id(tenant_id)
if not providers:
@@ -314,6 +280,7 @@ def list_tenant_added_models(tenant_id: str, model_type_filter: str = None):
return True, []
provider_instance_map: dict = {}
provider_info_map = {provider.id: provider for provider in providers}
instance_info_map = {instance.id: instance for instance in instances}
for provider_instance_record in instances:
provider_name = provider_info_map[provider_instance_record.provider_id].provider_name if provider_info_map.get(provider_instance_record.provider_id) else ""
if provider_instance_map.get(provider_name):
@@ -322,80 +289,17 @@ def list_tenant_added_models(tenant_id: str, model_type_filter: str = None):
provider_instance_map[provider_name] = [provider_instance_record]
model_records = TenantModelService.get_models_by_provider_ids_and_instance_ids(provider_ids, list({instance.id for instance in instances}))
target_type_records = [record for record in model_records if record.model_type == model_type_filter] if model_type_filter else model_records
model_record_map = {}
for model in target_type_records:
instance_model_key = f"{model.provider_id}|{model.instance_id}|{model.model_name}"
if model_record_map.get(instance_model_key):
model_record_map[instance_model_key].append(model)
else:
model_record_map[instance_model_key] = [model]
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
added_models = []
model_key_in_factory = []
provider_names = [provider.provider_name for provider in providers]
factory_rank_mapping = {factory["name"]: -_to_int(factory.get("rank", "500")) for factory in FACTORY_LLM_INFOS}
for factory in FACTORY_LLM_INFOS:
if factory["name"] not in provider_names:
continue
factory_instances = provider_instance_map.get(factory["name"])
if not factory_instances:
continue
for llm in factory["llm"]:
factory_model_types = _factory_model_types(llm)
if model_type_filter and model_type_filter not in factory_model_types:
continue
for factory_instance in factory_instances:
model_record_key = f"{factory_instance.provider_id}|{factory_instance.id}|{llm['llm_name']}"
model_key_in_factory.append(model_record_key)
manual_modified_models = model_record_map.get(model_record_key, [])
active_model_types = [manual_model.model_type for manual_model in manual_modified_models if manual_model.status == ActiveStatusEnum.ACTIVE.value]
inactive_model_types = [manual_model.model_type for manual_model in manual_modified_models if manual_model.status == ActiveStatusEnum.INACTIVE.value]
unsupport_model_types = [manual_model.model_type for manual_model in manual_modified_models if manual_model.status == ActiveStatusEnum.UNSUPPORTED.value]
model_types = list(set(factory_model_types + active_model_types) - set(inactive_model_types) - set(unsupport_model_types))
if not model_types:
continue
added_models.append(
{
"model_type": model_types,
"name": llm["llm_name"],
"provider_id": factory_instance.provider_id,
"provider_name": provider_info_map[factory_instance.provider_id].provider_name if provider_info_map.get(factory_instance.provider_id) else "",
"instance_id": factory_instance.id,
"instance_name": factory_instance.instance_name,
}
)
manual_added_model_record_keys = list(set(model_record_map.keys()) - set(model_key_in_factory))
if manual_added_model_record_keys:
instance_info_map = {instance.id: instance for instance in instances}
for model_record_key in manual_added_model_record_keys:
model_records = model_record_map.get(model_record_key, [])
if not model_records:
continue
# The internal key uses '|' as separator (UUID|UUID|model_name)
# since model_name may contain '@' characters.
try:
provider_id, instance_id, model_name = model_record_key.split("|", 2)
except ValueError:
logging.warning(f"Skipping malformed manual model record key: {model_record_key!r}")
continue
model_types = [model.model_type for model in model_records if model.status == ActiveStatusEnum.ACTIVE.value]
if not model_types:
continue
added_models.append(
{
"model_type": model_types,
"name": model_name,
"provider_id": provider_id,
"provider_name": provider_info_map[provider_id].provider_name if provider_info_map.get(provider_id) else "",
"instance_id": instance_id,
"instance_name": instance_info_map[instance_id].instance_name if instance_info_map.get(instance_id) else "",
}
)
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]
# Add TEI Builtin embedding model if configured
compose_profiles = os.getenv("COMPOSE_PROFILES", "")
@@ -415,6 +319,7 @@ def list_tenant_added_models(tenant_id: str, model_type_filter: str = None):
}
)
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"]))
return True, added_models

View File

@@ -18,13 +18,13 @@ import json
import logging
import asyncio
from common.constants import LLMType, ActiveStatusEnum
from common.misc_utils import get_uuid
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
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.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
from api.utils.model_utils import get_model_type_human, calculate_model_type
from rag.llm import ChatModel, EmbeddingModel, ModelMeta, OcrModel, RerankModel, TTSModel
@@ -66,7 +66,7 @@ def list_providers(tenant_id: str, all_available: bool = False):
List providers for a tenant.
If available_only is True, list all system-wide providers (pool providers).
Otherwise, list providers that the tenant has configured.
Otherwise, list providers that the tenant has configured, with a has_instance flag.
:param tenant_id: tenant ID
:param all_available: whether to list all available providers
@@ -102,11 +102,13 @@ def list_providers(tenant_id: str, all_available: bool = False):
for name in factory_names:
if name not in ["Youdao", "FastEmbed", "BAAI", "Builtin", "siliconflow_intl"] and factory_info_mapping.get(name):
factory_info = factory_info_mapping[name]
provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_name(tenant_id, name)
has_instance = bool(provider_obj and TenantModelInstanceService.get_all_by_provider_id(provider_obj.id))
model_types = sorted(set(model_type for llm in factory_info.get("llm", []) for model_type in _factory_model_types(llm))) if factory_info.get("llm", []) else []
if name in ["MinerU", "PaddleOCR", "OpenDataLoader"]:
model_types.append("ocr")
provider = {"model_types": model_types, "name": factory_info["name"], "url": {"default": factory_info.get("url", "")}}
provider = {"has_instance": has_instance, "model_types": model_types, "name": factory_info["name"], "url": {"default": factory_info.get("url", "")}}
if factory_info["name"].lower() == "siliconflow":
provider["url"]["intl"] = factory_info_map.get("siliconflow_intl", {}).get("url", "https://api.siliconflow.com/v1")
elif factory_info["name"] == "Tongyi-Qianwen":
@@ -199,7 +201,7 @@ async def list_provider_models(provider_id_or_name: str, api_key: str = None, ba
static_llms = [
{
"name": _factory_llm_name(llm),
"max_tokens": llm["max_tokens"],
"max_tokens": llm.get("max_tokens", 8192),
"model_types": _factory_model_types(llm),
"features": (llm.get("features") if llm.get("features") is not None else ((["is_tools"] if llm.get("is_tools") else []) + (["thinking"] if llm.get("thinking") else []))),
}
@@ -255,6 +257,181 @@ 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):
"""
Update a provider instance.
Updates the instance's api_key, base_url, region, and re-creates all models
based on the provided model_info list.
:param tenant_id: tenant ID
:param provider_id_or_name: provider/factory ID or name
:param instance_id_or_name: instance ID or name
:param instance_name: instance name (used as a logical identifier)
:param api_key: API key
:param base_url: base url
:param region: region
:param model_info: model info, [{
"model_type": ["chat"], # support multiple
"model_name": "name",
"max_tokens": 4096,
"extra": {
"is_tools": True
}
}]
:param verify: verify api_key
:return: (success, result_or_error_message)
"""
if not provider_id_or_name:
return False, "Provider ID or name is required"
provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_id(tenant_id, provider_id_or_name)
if not provider_obj:
provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_name(tenant_id, provider_id_or_name)
if not provider_obj:
return False, f"Provider '{provider_id_or_name}' does not exist"
provider_name = provider_obj.provider_name
# Find the instance
instance_obj = None
if instance_id_or_name:
_, instance_obj = TenantModelInstanceService.get_by_id(instance_id_or_name)
if instance_obj and instance_obj.provider_id != provider_obj.id:
instance_obj = None
if not instance_obj:
instance_obj = TenantModelInstanceService.get_by_provider_id_and_instance_name(provider_obj.id, instance_id_or_name)
if not instance_obj:
return False, f"No instance found for provider '{provider_id_or_name}' and instance '{instance_id_or_name}'"
base_url = _normalize_provider_base_url(provider_name, base_url)
api_key = _normalize_provider_api_key(provider_name, api_key)
api_key_str = ""
if api_key:
api_key_str = api_key if isinstance(api_key, str) else json.dumps(api_key)
# Verify api_key
model_verify_result = {}
if verify:
success, msg, model_verify_result = await verify_api_key(provider_name, api_key, base_url, region, model_info)
if not success:
return False, msg
# Update instance record
update_dict = {
"api_key": api_key_str,
}
if instance_name != instance_obj.instance_name:
update_dict["instance_name"] = instance_name
extra_fields = {}
if base_url:
extra_fields["base_url"] = base_url
if region:
extra_fields["region"] = region
# Preserve existing extra fields not overwritten
existing_extra = json.loads(instance_obj.extra) if instance_obj.extra else {}
existing_extra.update(extra_fields)
update_dict["extra"] = json.dumps(existing_extra)
TenantModelInstanceService.update_by_id(instance_obj.id, update_dict)
# Use the (possibly updated) instance_name for model recreation
effective_instance_name = instance_name
# Upsert models: add new ones, update existing ones, remove ones no longer selected
existing_model_objs = TenantModelService.get_models_by_instance_id(instance_obj.id)
existing_model_names = {model_obj.model_name: model_obj for model_obj in existing_model_objs}
# Delete models that are no longer in the submitted model_info
submitted_model_names = set()
if model_info:
submitted_model_names = {m.get("model_name") for m in model_info if m.get("model_name")}
elif model_info is not None:
# model_info is explicitly an empty list — remove all models
submitted_model_names = set()
models_to_remove = set(existing_model_names.keys()) - submitted_model_names
if models_to_remove:
TenantModelService.delete_by_ids([existing_model_names[n].id for n in models_to_remove])
msg = ""
if model_info:
for model in model_info:
model_name = model.get("model_name")
if not model_name:
continue
if verify:
verify_status = model_verify_result.get(model_name, ModelVerifyStatusEnum.UNKNOWN.value)
if model.get("extra"):
model["extra"].update({"verify": verify_status})
else:
model["extra"] = {"verify": verify_status}
if model_name in existing_model_names:
# Update existing model
update_dict = {}
if isinstance(model.get("model_type"), (str, list)):
target_model_type = calculate_model_type(model["model_type"])
if target_model_type != existing_model_names[model_name].model_type:
update_dict["model_type"] = target_model_type
merged_extra = json.loads(existing_model_names[model_name].extra) if existing_model_names[model_name].extra else {}
merged_extra.update(model["extra"])
if "max_tokens" in model:
merged_extra.update({"max_tokens": model["max_tokens"]})
update_dict["extra"] = json.dumps(merged_extra)
if update_dict:
TenantModelService.update_model(existing_model_names[model_name].id, update_dict)
else:
# Add new model
success, _msg = add_model_to_instance(tenant_id, provider_name, effective_instance_name, **model)
if not success:
msg += _msg
else:
if model_info is None:
# model_info not provided — add all factory default models (same as create)
factory_info = [f for f in FACTORY_LLM_INFOS if f["name"] == provider_name]
factory_llms = factory_info[0]["llm"]
for llm in factory_llms:
llm_name = _factory_llm_name(llm)
if llm_name in existing_model_names:
# Update existing
update_dict = {}
target_model_type = calculate_model_type(_factory_model_types(llm))
if target_model_type != existing_model_names[llm_name].model_type:
update_dict["model_type"] = target_model_type
db_extra = json.loads(existing_model_names[llm_name].extra) if existing_model_names[llm_name].extra else {}
db_extra_fields = {
"max_tokens": llm["max_tokens"],
"is_tools": llm.get("is_tools", False),
"thinking": "thinking" in llm.get("features", []),
}
if verify:
verify_status = model_verify_result.get(llm_name, ModelVerifyStatusEnum.UNKNOWN.value)
db_extra_fields["verify"] = verify_status
db_extra.update(db_extra_fields)
update_dict["extra"] = json.dumps(db_extra)
if update_dict:
TenantModelService.update_model(existing_model_names[llm_name].id, update_dict)
else:
extra_fields = {
"is_tools": llm.get("is_tools", False),
"thinking": "thinking" in llm.get("features", []),
}
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
})
if not success:
msg += _msg
return True, "success"
async def create_provider_instance(tenant_id: str, provider_id_or_name: str, instance_name: str, api_key: str | dict, base_url: str, region: str, model_info: list[dict] = None):
"""
Create a provider instance.
@@ -305,17 +482,9 @@ async def create_provider_instance(tenant_id: str, provider_id_or_name: str, ins
if api_key:
api_key_str = api_key if isinstance(api_key, str) else json.dumps(api_key)
# Only verify when there are models to probe. Generic providers such as
# "OpenAI-API-Compatible" may start empty and receive custom models later.
factory_entry = next((f for f in FACTORY_LLM_INFOS if f["name"] == provider_name), None)
if (factory_entry and factory_entry.get("llm")) or model_info:
success, msg = await verify_api_key(provider_name, api_key, base_url, region, model_info)
if not success:
return False, msg
success, msg = await verify_api_key(provider_name, api_key, base_url, region, model_info)
success, verify_msg, model_verify_result = await verify_api_key(provider_name, api_key, base_url, region, model_info)
if not success:
return False, msg
return False, verify_msg
extra_fields = {}
if base_url:
@@ -326,15 +495,72 @@ async def create_provider_instance(tenant_id: str, provider_id_or_name: str, ins
if model_info:
msg = ""
for model in model_info:
if model.get("extra"):
model["extra"].update({"verify": model_verify_result.get(model["model_name"], ModelVerifyStatusEnum.UNKNOWN.value)})
else:
model["extra"] = {"verify": model_verify_result.get(model["model_name"], ModelVerifyStatusEnum.UNKNOWN.value)}
success, _msg = add_model_to_instance(tenant_id, provider_name, instance_name, **model)
if not success:
msg += _msg
if msg:
return False, msg
else:
msg = ""
target_factory_name = "siliconflow_intl" if provider_name.lower() == "siliconflow" and region == "intl" else provider_name
factory_info = [f for f in FACTORY_LLM_INFOS if f["name"] == target_factory_name]
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)
}
})
if not success:
msg += _msg
if msg:
return False, msg
return True, "success"
async def create_name_only_provider_instance(tenant_id: str, provider_name: str, instance_name: str):
"""
Create a provider instance with only a name (no api_key/base_url validation).
:param tenant_id: tenant ID
:param provider_name: provider/factory name
:param instance_name: instance name (used as a logical identifier)
:return: (success, result_or_error_message)
"""
if not provider_name:
return False, "Provider name is required"
if instance_name == "default":
return False, "Instance name cannot be 'default'"
allowed_factories = [f["name"] for f in FACTORY_LLM_INFOS]
if provider_name not in allowed_factories:
return False, f"Provider '{provider_name}' is not allowed"
provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_name(tenant_id, provider_name)
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({})
)
return True, "success"
def list_provider_instances(tenant_id: str, provider_id_or_name: str):
"""
List provider instances for a tenant.
@@ -388,7 +614,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
:return: (success, result_or_error_message)
"""
if not provider_id_or_name:
return False, "Provider ID or name is required"
return False, "Provider ID or name is required", {}
provider_obj = None
if provider_id_or_name:
@@ -405,24 +631,21 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
factory_info = [f for f in FACTORY_LLM_INFOS if f["name"] == target_factory_name]
if not factory_info:
return False, f"Provider '{provider_id_or_name}' not found"
return False, f"Provider '{provider_id_or_name}' not found", {}
factory_llms = factory_info[0]["llm"]
if not factory_llms:
if not model_info:
return False, f"No models found for provider '{provider_id_or_name}'"
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 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", [])]
if not factory_llms:
return False, f"No valid models found for provider '{provider_id_or_name}'"
return False, f"No valid models found for provider '{provider_id_or_name}'", {}
else:
factory_llms = factory_info[0]["llm"]
if not factory_llms:
return False, f"No models found for provider '{provider_id_or_name}'", {}
model_verify_result = {}
# test if api key works
chat_passed, embd_passed, rerank_passed, ocr_passed, tts_passed = False, False, False, False, False
timeout_seconds = int(os.environ.get("LLM_TIMEOUT_SECONDS", 10))
@@ -448,6 +671,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
if len(arr[0]) == 0:
raise Exception("Fail")
embd_passed = True
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.SUCCESS.value
except Exception as e:
logging.exception(
"Fail to access embedding model for provider=%s model=%s",
@@ -455,6 +679,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
llm["llm_name"],
)
msg += f"\nFail to access embedding model({llm['llm_name']}) using this api key." + str(e)
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
elif not chat_passed and LLMType.CHAT.value in model_types:
assert provider_name in ChatModel, f"Chat model from {provider_name} is not supported yet."
mdl = ChatModel[provider_name](api_key_str, llm["llm_name"], base_url=base_url, **extra)
@@ -472,8 +697,10 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
result = await asyncio.wait_for(check_streamly(), timeout=timeout_seconds)
if result:
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.SUCCESS.value
chat_passed = True
else:
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
raise Exception("No valid response received")
except Exception as e:
logging.exception(
@@ -481,6 +708,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
provider_name,
llm["llm_name"],
)
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
msg += f"\nFail to access model({provider_name}/{llm['llm_name']}) using this api key." + str(e)
elif not rerank_passed and LLMType.RERANK.value in model_types:
if provider_name not in RerankModel:
@@ -497,6 +725,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
if len(arr) == 0 or tc == 0:
raise Exception("Fail")
rerank_passed = True
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.SUCCESS.value
logging.debug(f"passed model rerank {llm['llm_name']}")
except Exception as e:
logging.exception(
@@ -504,6 +733,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
provider_name,
llm["llm_name"],
)
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
msg += f"\nFail to access model({provider_name}/{llm['llm_name']}) using this api key." + str(e)
elif not ocr_passed and LLMType.OCR.value in model_types:
assert provider_name in OcrModel, f"OCR model from {provider_name} is not supported yet."
@@ -515,6 +745,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
)
if not ok:
raise RuntimeError(reason or "Model not available")
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.SUCCESS.value
ocr_passed = True
except Exception as e:
logging.exception(
@@ -522,6 +753,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
provider_name,
llm["llm_name"],
)
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
msg += f"\nFail to access model({provider_name}/{llm['llm_name']})." + str(e)
elif not tts_passed and LLMType.TTS.value in model_types:
assert provider_name in TTSModel, f"TTS model from {provider_name} is not supported yet."
@@ -536,6 +768,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
asyncio.to_thread(drain_tts),
timeout=timeout_seconds,
)
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.SUCCESS.value
tts_passed = True
except Exception as e:
logging.exception(
@@ -543,13 +776,14 @@ async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url
provider_name,
llm["llm_name"],
)
model_verify_result[llm["llm_name"]] = ModelVerifyStatusEnum.FAIL.value
msg += f"\nFail to access model({provider_name}/{llm['llm_name']})." + str(e)
if any([embd_passed, chat_passed, rerank_passed, ocr_passed, tts_passed]):
msg = ""
break
success = any([embd_passed, chat_passed, rerank_passed, ocr_passed, tts_passed])
return success, "success" if success else msg
return success, "success" if success else msg, model_verify_result
def show_provider_instance(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str):
@@ -578,7 +812,16 @@ def show_provider_instance(tenant_id: str, provider_id_or_name: str, instance_id
return False, f"No instance found for provider '{provider_id_or_name}' and instance '{instance_id_or_name}'"
extra_fields = json.loads(instance_obj.extra) if instance_obj.extra else {}
return True, {"id": instance_obj.id, "instance_name": instance_obj.instance_name, "provider_id": provider_id, "region": extra_fields.get("region", ""), "status": instance_obj.status}
return True, {
"id": instance_obj.id,
"instance_name": instance_obj.instance_name,
"provider_id": provider_id,
"region": extra_fields.get("region", ""),
"base_url": extra_fields.get("base_url", ""),
"api_key": instance_obj.api_key,
"status": instance_obj.status
}
def drop_provider_instances(tenant_id: str, provider_id_or_name: str, instance_id_or_names: list):
@@ -618,55 +861,6 @@ def drop_provider_instances(tenant_id: str, provider_id_or_name: str, instance_i
return True, None
def _hybrid_get_instance_models(provider_name: str, instance_id: str):
# List all models from the LLM dictionary for this provider
factory_info = [f for f in FACTORY_LLM_INFOS if f["name"] == provider_name]
if not factory_info:
return False, f"Provider '{provider_name}' not found"
# Get model records for this instance from tenant_model table
model_records = TenantModelService.get_models_by_instance_id(instance_id)
# Build a map of model_name -> status, type
model_info_map: dict = {}
model_unsupported_type_map = {}
for model_record in model_records:
if model_record.status == ActiveStatusEnum.UNSUPPORTED.value:
if model_unsupported_type_map.get(model_record.model_name):
model_unsupported_type_map[model_record.model_name].append(model_record.model_type)
else:
model_unsupported_type_map[model_record.model_name] = [model_record.model_type]
continue
if model_info_map.get(model_record.model_name):
model_info_map[model_record.model_name]["model_type"].append(model_record.model_type)
else:
model_info_map[model_record.model_name] = {"status": model_record.status, "model_type": [model_record.model_type], "extra": model_record.extra}
llms = factory_info[0].get("llm", [])
models = []
for llm in llms:
models.append(
{
"name": llm["llm_name"],
"model_type": list(set(_factory_model_types(llm) + model_info_map.get(llm["llm_name"], {}).get("model_type", [])) - set(model_unsupported_type_map.get(llm["llm_name"], []))),
"max_tokens": llm.get("max_tokens"),
"status": model_info_map.get(llm["llm_name"], {}).get("status", "active"),
}
)
factory_models = [m["name"] for m in models]
for model_name, model_info_dict in model_info_map.items():
if model_name not in factory_models:
extra_fields = json.loads(model_info_dict["extra"]) if model_info_dict["extra"] else {}
models.append(
{
"name": model_name,
"model_type": set(model_info_dict["model_type"]) - set(model_unsupported_type_map.get(model_name, [])),
"max_tokens": extra_fields.get("max_tokens", 8192),
"status": model_info_dict["status"],
}
)
return True, models
def list_instance_models(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, supported_only: bool = False):
"""
List models for a provider instance.
@@ -708,8 +902,21 @@ def list_instance_models(tenant_id: str, provider_id_or_name: str, instance_id_o
instance_obj = TenantModelInstanceService.get_by_provider_id_and_instance_name(provider_obj.id, instance_id_or_name)
if not instance_obj:
return False, f"No instance found for provider '{provider_id_or_name}' and instance '{instance_id_or_name}'"
# Get models
model_objs = TenantModelService.get_models_by_instance_id(instance_obj.id)
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 [])
})
return _hybrid_get_instance_models(provider_obj.provider_name, instance_obj.id)
return True, model_list
def update_instance_models(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, model_names: list, model_types: list):
@@ -731,19 +938,16 @@ def update_instance_models(tenant_id: str, provider_id_or_name: str, instance_id
if not instance_obj:
return False, f"No instance found for provider '{provider_id_or_name}' and instance '{instance_id_or_name}'"
found, models = _hybrid_get_instance_models(provider_obj.provider_name, instance_obj.id)
if not found:
return False, models
model_info_map = {model["name"]: model for model in models}
not_exist_models = set(model_names) - set(model_info_map.keys())
model_objs = TenantModelService.get_models_by_instance_id(instance_obj.id)
not_exist_models = set(model_names) - {model_obj.model_name for model_obj in model_objs}
if not_exist_models:
return False, f"Models {not_exist_models} not found for provider '{provider_id_or_name}' and instance '{instance_id_or_name}'"
for model_name in model_names:
model_info = model_info_map.get(model_name, {})
TenantModelService.upsert_model_type(
provider_obj.id, instance_obj.id, model_name, {"add": list(set(model_types) - set(model_info["model_type"])), "delete": list(set(model_info["model_type"]) - set(model_types))}
)
target_model_type_bin = calculate_model_type(model_types)
to_update = [model_obj.id for model_obj in model_objs if model_obj.model_type != target_model_type_bin and model_obj.model_name in model_names]
if to_update:
TenantModelService.batch_update_model_type(to_update, target_model_type_bin)
return True, "success"
@@ -772,22 +976,26 @@ def add_model_to_instance(tenant_id: str, provider_id_or_name: str, instance_id_
if isinstance(model_type, str):
model_type = [model_type]
for _type in model_type:
extra_fields = {"max_tokens": max_tokens}
target_model = [llm for llm in llms if _type in _factory_model_types(llm) and llm["llm_name"] == model_name]
if target_model:
extra_fields.update({"is_tools": target_model[0].get("is_tools", False)})
if extra:
if provider_id_or_name == "SoMark" and LLMType.OCR.value in model_type:
extra_fields["ocr_config"] = extra
else:
extra_fields.update(extra)
TenantModelService.insert(model_name=model_name, provider_id=provider_obj.id, instance_id=instance_obj.id, model_type=_type, extra=json.dumps(extra_fields))
model_type_bin = calculate_model_type(model_type)
extra_fields = {"max_tokens": max_tokens}
target_model = [llm for llm in llms if llm["llm_name"] == model_name]
if target_model:
extra_fields.update({"is_tools": target_model[0].get("is_tools", False)})
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)
)
return True, "success"
def update_model_status(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, model_name: str, status: str):
def update_model(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, model_name: str, update_dict: dict):
"""
Enable or disable a model for a provider instance.
@@ -800,10 +1008,12 @@ def update_model_status(tenant_id: str, provider_id_or_name: str, instance_id_or
:param provider_id_or_name: provider/factory ID or name
:param instance_id_or_name: instance ID or name
:param model_name: model name
:param status: "active" or "inactive" (ActiveStatusEnum values)
:param update_dict:
status: "active" or "inactive" (ActiveStatusEnum values)
max_tokens: > 0
:return: (success, result_or_error_message)
"""
if status not in (ActiveStatusEnum.ACTIVE.value, ActiveStatusEnum.INACTIVE.value):
if update_dict.get("status") and update_dict["status"] not in (ActiveStatusEnum.ACTIVE.value, ActiveStatusEnum.INACTIVE.value):
return False, f"status must be '{ActiveStatusEnum.ACTIVE.value}' or '{ActiveStatusEnum.INACTIVE.value}'"
# Check if provider exists for this tenant
@@ -824,39 +1034,65 @@ def update_model_status(tenant_id: str, provider_id_or_name: str, instance_id_or
if not instance_obj:
return False, f"No instance found for provider '{provider_id_or_name}' and instance '{instance_id_or_name}'"
# Check if model record already exists in tenant_model table
model_obj_list = TenantModelService.get_by_provider_id_and_instance_id_and_model_name(provider_obj.id, instance_obj.id, model_name)
model_obj = TenantModelService.get_by_provider_id_and_instance_id_and_model_name(provider_obj.id, instance_obj.id, model_name)
to_update = {}
if "status" in update_dict and update_dict.get("status") != model_obj.status:
to_update.update({"status": update_dict["status"]})
new_extra = update_dict.get("extra", {})
if "max_tokens" in update_dict:
new_extra.update({"max_tokens": update_dict["max_tokens"]})
if "verify" in update_dict:
new_extra.update({"verify": update_dict["verify"]})
if new_extra:
db_extra = json.loads(model_obj.extra)
db_extra.update(**new_extra)
to_update.update({"extra": json.dumps(db_extra)})
if "model_type" in update_dict:
target_model_type = calculate_model_type(update_dict["model_type"])
if target_model_type != model_obj.model_type:
to_update.update({"model_type": target_model_type})
if model_obj_list:
# Model record exists — update its status
TenantModelService.batch_update_model_status([m.id for m in model_obj_list if m.status != ActiveStatusEnum.UNSUPPORTED.value], status)
else:
# Model record does not exist
if status == ActiveStatusEnum.ACTIVE.value:
# Default is active, no need to add a record
return True, None
# status is "inactive" — create a record with inactive status
# Look up model schema from FACTORY_LLM_INFOS
factory_info = [f for f in FACTORY_LLM_INFOS if f["name"] == provider_obj.provider_name]
if not factory_info:
return False, f"Provider '{provider_id_or_name}' not found"
llms = factory_info[0].get("llm", [])
target_llm = [llm for llm in llms if llm["llm_name"] == model_name]
if not target_llm:
return False, f"provider {provider_obj.provider_name} model {model_name} not found"
if to_update:
TenantModelService.update_model(model_obj.id, to_update)
for model_type in _factory_model_types(target_llm[0]):
TenantModelService.insert(
id=get_uuid(),
model_name=model_name,
model_type=model_type,
provider_id=provider_obj.id,
instance_id=instance_obj.id,
status=status,
extra=json.dumps({"max_tokens": target_llm[0].get("max_tokens", 8192), "is_tools": target_llm[0].get("is_tools", False)}),
)
return True, "success"
return True, None
async def delete_models_from_instance(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, model_name: list[str]):
"""
Delete models from instance.
:param tenant_id: tenant ID
:param provider_id_or_name: provider/factory ID or name
:param instance_id_or_name: instance ID or name
:param model_name: list of model name
"""
# Check if provider exists for this tenant (by ID first, then by name)
provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_id(tenant_id, provider_id_or_name)
if not provider_obj:
provider_obj = TenantModelProviderService.get_by_tenant_id_and_provider_name(tenant_id, provider_id_or_name)
if not provider_obj:
return False, f"No provider found for provider '{provider_id_or_name}'"
# Check if instance exists (by ID first, then by name)
instance_obj = None
if instance_id_or_name:
_, instance_obj = TenantModelInstanceService.get_by_id(instance_id_or_name)
if instance_obj and instance_obj.provider_id != provider_obj.id:
instance_obj = None
if not instance_obj:
instance_obj = TenantModelInstanceService.get_by_provider_id_and_instance_name(provider_obj.id, instance_id_or_name)
if not instance_obj:
return False, f"No instance found for provider '{provider_id_or_name}' and instance '{instance_id_or_name}'"
model_objs = TenantModelService.get_models_by_instance_id(instance_obj.id)
not_exist_models = set(model_name) - {model_obj.model_name for model_obj in model_objs}
if not_exist_models:
return False, f"Models {not_exist_models} not found for provider '{provider_id_or_name}' and instance '{instance_id_or_name}'"
TenantModelService.delete_by_ids([model_obj.id for model_obj in model_objs if model_obj.model_name in model_name])
return True, "success"
async def chat_to_model(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, model_name: str, message: str, stream: bool = False, thinking: bool = False):
@@ -896,7 +1132,7 @@ async def chat_to_model(tenant_id: str, provider_id_or_name: str, instance_id_or
# Get model config
composite_name = f"{model_name}@{instance_name}@{provider_name}"
try:
model_config = get_model_config_from_provider_instance(tenant_id, LLMType.CHAT.value, composite_name)
model_config = get_model_config_from_provider_instance(tenant_id, LLMType.CHAT, composite_name)
except LookupError:
return False, f"Model '{composite_name}' not authorized"