mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-08 08:28:02 +08:00
Feat: SoMark (#16482)
Follow #15486 Co-authored-by: limuting <limuting233@gmail.com> Co-authored-by: lutianyi <lutianyi233@163.com> Co-authored-by: justinychuang <huangyicheng@soulcode.cn> Co-authored-by: maybehokori <138367708+maybehokori@users.noreply.github.com>
This commit is contained in:
@@ -82,20 +82,10 @@ def list_providers(tenant_id: str, all_available: bool = False):
|
||||
for factory_info in FACTORY_LLM_INFOS:
|
||||
if factory_info["name"] in ["Youdao", "FastEmbed", "BAAI", "Builtin", "siliconflow_intl"]:
|
||||
continue
|
||||
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 []
|
||||
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 factory_info["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 = {"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":
|
||||
@@ -112,21 +102,11 @@ 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]
|
||||
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 []
|
||||
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 = {"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":
|
||||
@@ -155,10 +135,7 @@ def add_provider(tenant_id: str, provider_name: str):
|
||||
if existing:
|
||||
return False, f"Provider {provider_name} already exists"
|
||||
|
||||
TenantModelProviderService.insert(
|
||||
tenant_id=tenant_id,
|
||||
provider_name=provider_name
|
||||
)
|
||||
TenantModelProviderService.insert(tenant_id=tenant_id, provider_name=provider_name)
|
||||
return True, "success"
|
||||
|
||||
|
||||
@@ -195,17 +172,11 @@ def show_provider(provider_id_or_name: str):
|
||||
if provider_id_or_name:
|
||||
_, provider_obj = TenantModelProviderService.get_by_id(provider_id_or_name)
|
||||
provider_name = provider_obj.provider_name if provider_obj else provider_id_or_name
|
||||
fac_list = [f for f in FACTORY_LLM_INFOS if f["name"]==provider_name]
|
||||
fac_list = [f for f in FACTORY_LLM_INFOS if f["name"] == provider_name]
|
||||
if not fac_list:
|
||||
return False, f"Provider '{provider_id_or_name}' not found"
|
||||
factory_info = fac_list[0]
|
||||
return True, {
|
||||
"base_url": {
|
||||
"default": factory_info.get("url", "")
|
||||
},
|
||||
"name": factory_info["name"],
|
||||
"total_models": len(factory_info.get("llm", []))
|
||||
}
|
||||
return True, {"base_url": {"default": factory_info.get("url", "")}, "name": factory_info["name"], "total_models": len(factory_info.get("llm", []))}
|
||||
|
||||
|
||||
async def list_provider_models(provider_id_or_name: str, api_key: str = None, base_url: str = None):
|
||||
@@ -225,19 +196,15 @@ async def list_provider_models(provider_id_or_name: str, api_key: str = None, ba
|
||||
if not factory_info:
|
||||
return False, f"Provider '{provider_id_or_name}' not found"
|
||||
api_key = _normalize_provider_api_key(provider_name, api_key)
|
||||
static_llms = [{
|
||||
static_llms = [
|
||||
{
|
||||
"name": _factory_llm_name(llm),
|
||||
"max_tokens": llm["max_tokens"],
|
||||
"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 [])
|
||||
)
|
||||
)
|
||||
} for llm in factory_info[0]["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 []))),
|
||||
}
|
||||
for llm in factory_info[0]["llm"]
|
||||
]
|
||||
|
||||
model_base_url = _normalize_provider_base_url(provider_name, base_url) or factory_info[0].get("url", "")
|
||||
remote_models = []
|
||||
@@ -284,11 +251,11 @@ def show_provider_model(provider_id_or_name: str, model_name: str):
|
||||
"max_tokens": llm_info["max_tokens"],
|
||||
"model_types": _factory_model_types(llm_info),
|
||||
"thinking": None,
|
||||
"model_type_map": {model_type: True for model_type in _factory_model_types(llm_info)}
|
||||
"model_type_map": {model_type: True for model_type in _factory_model_types(llm_info)},
|
||||
}
|
||||
|
||||
|
||||
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):
|
||||
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.
|
||||
|
||||
@@ -346,12 +313,35 @@ async def create_provider_instance(tenant_id: str, provider_id_or_name: str, ins
|
||||
if not success:
|
||||
return False, msg
|
||||
|
||||
# For SoMark, embed OCR config from model_info into the api_key JSON so
|
||||
# SoMarkOcrModel.__init__ can read it via the existing key → api_key_payload
|
||||
# path. This avoids changing the deprecated LLMBundle in tenant_llm_service.py.
|
||||
if provider_name == "SoMark" and model_info:
|
||||
cfg = {}
|
||||
if api_key_str:
|
||||
try:
|
||||
cfg = json.loads(api_key_str)
|
||||
except Exception:
|
||||
pass
|
||||
if not isinstance(cfg, dict):
|
||||
cfg = {}
|
||||
for model in model_info:
|
||||
if model.get("extra"):
|
||||
cfg.update(model["extra"])
|
||||
if base_url:
|
||||
cfg["SOMARK_BASE_URL"] = base_url
|
||||
api_key_str = json.dumps(cfg)
|
||||
|
||||
success, msg = await verify_api_key(provider_name, api_key, base_url, region, model_info)
|
||||
if not success:
|
||||
return False, msg
|
||||
|
||||
extra_fields = {}
|
||||
if base_url:
|
||||
extra_fields["base_url"] = base_url
|
||||
if region:
|
||||
extra_fields["region"] = region
|
||||
TenantModelInstanceService.create_instance(provider_id=provider_obj.id,instance_name=instance_name,api_key=api_key_str, extra=json.dumps(extra_fields))
|
||||
TenantModelInstanceService.create_instance(provider_id=provider_obj.id, instance_name=instance_name, api_key=api_key_str, extra=json.dumps(extra_fields))
|
||||
if model_info:
|
||||
msg = ""
|
||||
for model in model_info:
|
||||
@@ -384,18 +374,20 @@ def list_provider_instances(tenant_id: str, provider_id_or_name: str):
|
||||
instances = []
|
||||
for instance_obj in instance_objs:
|
||||
extra_fields = json.loads(instance_obj.extra) if instance_obj.extra else {}
|
||||
instances.append({
|
||||
"id": instance_obj.id,
|
||||
"instance_name": instance_obj.instance_name,
|
||||
"provider_id": provider_id,
|
||||
"region": extra_fields.get("region", ""),
|
||||
"status": instance_obj.status,
|
||||
})
|
||||
instances.append(
|
||||
{
|
||||
"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, instances
|
||||
|
||||
|
||||
async def verify_api_key(provider_id_or_name: str, api_key: str|dict, base_url: str=None, region: str=None, model_info: list[dict]=None):
|
||||
async def verify_api_key(provider_id_or_name: str, api_key: str | dict, base_url: str = None, region: str = None, model_info: list[dict] = None):
|
||||
"""
|
||||
Verify API key for a provider.
|
||||
|
||||
@@ -438,10 +430,15 @@ async def verify_api_key(provider_id_or_name: str, api_key: str|dict, base_url:
|
||||
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", []) ]
|
||||
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}'"
|
||||
|
||||
@@ -481,11 +478,12 @@ async def verify_api_key(provider_id_or_name: str, api_key: str|dict, base_url:
|
||||
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)
|
||||
try:
|
||||
|
||||
async def check_streamly():
|
||||
async for chunk in mdl.async_chat_streamly(
|
||||
None,
|
||||
[{"role": "user", "content": "Hi"}],
|
||||
{"temperature": 0.9},
|
||||
None,
|
||||
[{"role": "user", "content": "Hi"}],
|
||||
{"temperature": 0.9},
|
||||
):
|
||||
if chunk and isinstance(chunk, str) and chunk.find("**ERROR**") < 0:
|
||||
return True
|
||||
@@ -548,6 +546,7 @@ async def verify_api_key(provider_id_or_name: str, api_key: str|dict, base_url:
|
||||
assert provider_name in TTSModel, f"TTS model from {provider_name} is not supported yet."
|
||||
mdl = TTSModel[provider_name](key=api_key_str, model_name=llm["llm_name"], base_url=base_url)
|
||||
try:
|
||||
|
||||
def drain_tts():
|
||||
for _ in mdl.tts("Hello~ RAGFlower!"):
|
||||
pass
|
||||
@@ -598,13 +597,7 @@ 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", ""), "status": instance_obj.status}
|
||||
|
||||
|
||||
def drop_provider_instances(tenant_id: str, provider_id_or_name: str, instance_id_or_names: list):
|
||||
@@ -665,33 +658,31 @@ def _hybrid_get_instance_models(provider_name: str, instance_id: str):
|
||||
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
|
||||
}
|
||||
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"),
|
||||
})
|
||||
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"],
|
||||
})
|
||||
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
|
||||
|
||||
|
||||
@@ -770,18 +761,12 @@ def update_instance_models(tenant_id: str, provider_id_or_name: str, instance_id
|
||||
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))
|
||||
}
|
||||
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))}
|
||||
)
|
||||
return True, "success"
|
||||
|
||||
|
||||
def add_model_to_instance(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, model_name: str, model_type: str|list[str], max_tokens: int=8192, extra: dict=None):
|
||||
def add_model_to_instance(tenant_id: str, provider_id_or_name: str, instance_id_or_name: str, model_name: str, model_type: str | list[str], max_tokens: int = 8192, extra: dict = None):
|
||||
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)
|
||||
@@ -812,14 +797,11 @@ def add_model_to_instance(tenant_id: str, provider_id_or_name: str, instance_id_
|
||||
if target_model:
|
||||
extra_fields.update({"is_tools": target_model[0].get("is_tools", False)})
|
||||
if extra:
|
||||
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)
|
||||
)
|
||||
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))
|
||||
|
||||
return True, "success"
|
||||
|
||||
@@ -862,9 +844,7 @@ def update_model_status(tenant_id: str, provider_id_or_name: str, instance_id_or
|
||||
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_list = TenantModelService.get_by_provider_id_and_instance_id_and_model_name(provider_obj.id, instance_obj.id, model_name)
|
||||
|
||||
if model_obj_list:
|
||||
# Model record exists — update its status
|
||||
@@ -892,7 +872,7 @@ def update_model_status(tenant_id: str, provider_id_or_name: str, instance_id_or
|
||||
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)})
|
||||
extra=json.dumps({"max_tokens": target_llm[0].get("max_tokens", 8192), "is_tools": target_llm[0].get("is_tools", False)}),
|
||||
)
|
||||
|
||||
return True, None
|
||||
|
||||
@@ -18,7 +18,18 @@ import os
|
||||
import enum
|
||||
import json
|
||||
from common import settings
|
||||
from common.constants import ActiveStatusEnum, LLMType, MINERU_DEFAULT_CONFIG, MINERU_ENV_KEYS, OPENDATALOADER_DEFAULT_CONFIG, OPENDATALOADER_ENV_KEYS, PADDLEOCR_DEFAULT_CONFIG, PADDLEOCR_ENV_KEYS
|
||||
from common.constants import (
|
||||
ActiveStatusEnum,
|
||||
LLMType,
|
||||
MINERU_DEFAULT_CONFIG,
|
||||
MINERU_ENV_KEYS,
|
||||
OPENDATALOADER_DEFAULT_CONFIG,
|
||||
OPENDATALOADER_ENV_KEYS,
|
||||
PADDLEOCR_DEFAULT_CONFIG,
|
||||
PADDLEOCR_ENV_KEYS,
|
||||
SOMARK_DEFAULT_CONFIG,
|
||||
SOMARK_ENV_KEYS,
|
||||
)
|
||||
from api.db.services.tenant_llm_service import TenantService
|
||||
from api.db.services.tenant_model_provider_service import TenantModelProviderService
|
||||
from api.db.services.tenant_model_instance_service import TenantModelInstanceService
|
||||
@@ -247,6 +258,11 @@ def get_model_config_from_provider_instance(tenant_id, model_type: str | enum.En
|
||||
"is_tools": model_extra.get("is_tools", is_tool),
|
||||
"max_tokens": model_extra.get("max_tokens") or 8192,
|
||||
}
|
||||
if provider_name.lower() == "somark":
|
||||
# SoMark/OCR factories read parser config (somark_*, parse_method, ...)
|
||||
# from model_config["extra"]; see tenant_llm_service.LLMBundle OCR path.
|
||||
model_config["extra"] = model_extra.get("ocr_config", model_extra)
|
||||
|
||||
if api_key_payload is not None:
|
||||
model_config["api_key_payload"] = api_key_payload
|
||||
|
||||
@@ -337,6 +353,15 @@ def ensure_opendataloader_from_env(tenant_id: str) -> str | None:
|
||||
)
|
||||
|
||||
|
||||
def ensure_somark_from_env(tenant_id: str) -> str | None:
|
||||
return _ensure_ocr_provider_from_env(
|
||||
tenant_id,
|
||||
"SoMark",
|
||||
"somark-from-env",
|
||||
_collect_env_config(SOMARK_ENV_KEYS, SOMARK_DEFAULT_CONFIG),
|
||||
)
|
||||
|
||||
|
||||
def get_models_by_tenant_and_provider_and_model_type(tenant_id: str, provider_name: str, model_type: str):
|
||||
"""
|
||||
Query TenantModel records by tenant_id, provider_name and model_name.
|
||||
|
||||
Reference in New Issue
Block a user