mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-25 09:53:29 +08:00
Feat: v0.27.0 model provider (#16604)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user