Feat/tenant model (#13072)

### What problem does this PR solve?

Add id for table tenant_llm and apply in LLMBundle.

### Type of change

- [x] Refactoring

---------

Co-authored-by: Yingfeng <yingfeng.zhang@gmail.com>
Co-authored-by: Liu An <asiro@qq.com>
This commit is contained in:
Lynn
2026-03-05 17:27:17 +08:00
committed by GitHub
parent 47540a4147
commit 62cb292635
54 changed files with 1754 additions and 361 deletions
@@ -121,6 +121,132 @@ def _load_session_module(monkeypatch):
common_pkg.__path__ = [str(repo_root / "common")]
monkeypatch.setitem(sys.modules, "common", common_pkg)
# Mock common.constants module
from enum import Enum
from strenum import StrEnum
class _StubLLMType(StrEnum):
CHAT = "chat"
EMBEDDING = "embedding"
SPEECH2TEXT = "speech2text"
IMAGE2TEXT = "image2text"
RERANK = "rerank"
TTS = "tts"
OCR = "ocr"
class _StubParserType(StrEnum):
PRESENTATION = "presentation"
LAWS = "laws"
MANUAL = "manual"
PAPER = "paper"
RESUME = "resume"
BOOK = "book"
QA = "qa"
TABLE = "table"
NAIVE = "naive"
PICTURE = "picture"
ONE = "one"
AUDIO = "audio"
EMAIL = "email"
KG = "knowledge_graph"
TAG = "tag"
class _StubRetCode(int, Enum):
SUCCESS = 0
NOT_EFFECTIVE = 10
EXCEPTION_ERROR = 100
ARGUMENT_ERROR = 101
DATA_ERROR = 102
OPERATING_ERROR = 103
CONNECTION_ERROR = 105
RUNNING = 106
PERMISSION_ERROR = 108
AUTHENTICATION_ERROR = 109
BAD_REQUEST = 400
UNAUTHORIZED = 401
SERVER_ERROR = 500
FORBIDDEN = 403
NOT_FOUND = 404
CONFLICT = 409
class _StubStatusEnum(str, Enum):
VALID = "1"
INVALID = "0"
class _StubActiveEnum(Enum):
ACTIVE = "1"
INACTIVE = "0"
class _StubStorage(Enum):
MINIO = 1
AZURE_SPN = 2
AZURE_SAS = 3
AWS_S3 = 4
OSS = 5
OPENDAL = 6
GCS = 7
class _StubMCPServerType(StrEnum):
SSE = "sse"
STREAMABLE_HTTP = "streamable-http"
class _StubTaskStatus(StrEnum):
UNSTART = "0"
RUNNING = "1"
CANCEL = "2"
DONE = "3"
FAIL = "4"
SCHEDULE = "5"
class _StubFileSource(StrEnum):
LOCAL = ""
KNOWLEDGEBASE = "knowledgebase"
S3 = "s3"
NOTION = "notion"
DISCORD = "discord"
CONFLUENCE = "confluence"
GMAIL = "gmail"
GOOGLE_DRIVE = "google_drive"
JIRA = "jira"
SHAREPOINT = "sharepoint"
SLACK = "slack"
TEAMS = "teams"
WEBDAV = "webdav"
MOODLE = "moodle"
DROPBOX = "dropbox"
BOX = "box"
R2 = "r2"
OCI_STORAGE = "oci_storage"
GOOGLE_CLOUD_STORAGE = "google_cloud_storage"
AIRTABLE = "airtable"
ASANA = "asana"
GITHUB = "github"
GITLAB = "gitlab"
IMAP = "imap"
BITBUCKET = "bitbucket"
ZENDESK = "zendesk"
SEAFILE = "seafile"
MYSQL = "mysql"
POSTGRESQL = "postgresql"
common_constants_mod = ModuleType("common.constants")
common_constants_mod.LLMType = _StubLLMType
common_constants_mod.ParserType = _StubParserType
common_constants_mod.RetCode = _StubRetCode
common_constants_mod.StatusEnum = _StubStatusEnum
common_constants_mod.ActiveEnum = _StubActiveEnum
common_constants_mod.Storage = _StubStorage
common_constants_mod.MCPServerType = _StubMCPServerType
common_constants_mod.TaskStatus = _StubTaskStatus
common_constants_mod.FileSource = _StubFileSource
common_constants_mod.SERVICE_CONF = "service_conf.yaml"
common_constants_mod.RAG_FLOW_SERVICE_NAME = "ragflow"
common_constants_mod.SVR_QUEUE_NAME = "rag_flow_svr_queue"
common_constants_mod.SVR_CONSUMER_GROUP_NAME = "rag_flow_svr_task_broker"
common_constants_mod.PAGERANK_FLD = "pagerank_fea"
common_constants_mod.TAG_FLD = "tag_feas"
monkeypatch.setitem(sys.modules, "common.constants", common_constants_mod)
deepdoc_pkg = ModuleType("deepdoc")
deepdoc_parser_pkg = ModuleType("deepdoc.parser")
deepdoc_parser_pkg.__path__ = []
@@ -166,6 +292,180 @@ def _load_session_module(monkeypatch):
monkeypatch.setitem(sys.modules, "deepdoc.parser.utils", deepdoc_parser_utils)
monkeypatch.setitem(sys.modules, "xgboost", ModuleType("xgboost"))
# Mock tenant_llm_service for TenantLLMService and TenantService
tenant_llm_service_mod = ModuleType("api.db.services.tenant_llm_service")
class _MockModelConfig:
def __init__(self, tenant_id, model_name):
self.tenant_id = tenant_id
self.llm_name = model_name
self.llm_factory = "Builtin"
self.api_key = "fake-api-key"
self.api_base = "https://api.example.com"
self.model_type = "chat"
self.max_tokens = 8192
self.used_tokens = 0
self.status = 1
self.id = 1
def to_dict(self):
return {
"tenant_id": self.tenant_id,
"llm_name": self.llm_name,
"llm_factory": self.llm_factory,
"api_key": self.api_key,
"api_base": self.api_base,
"model_type": self.model_type,
"max_tokens": self.max_tokens,
"used_tokens": self.used_tokens,
"status": self.status,
"id": self.id
}
class _StubTenantService:
@staticmethod
def get_by_id(tenant_id):
# Return a mock tenant with default model configurations
return True, SimpleNamespace(
id=tenant_id,
llm_id="chat-model",
embd_id="embd-model",
asr_id="asr-model",
img2txt_id="img2txt-model",
rerank_id="rerank-model",
tts_id="tts-model"
)
class _StubTenantLLMService:
@staticmethod
def get_api_key(tenant_id, model_name):
return _MockModelConfig(tenant_id, model_name)
@staticmethod
def split_model_name_and_factory(model_name):
if "@" in model_name:
parts = model_name.split("@")
return parts[0], parts[1]
return model_name, None
class _StubLLMFactoriesService:
@staticmethod
def query(**_kwargs):
return []
tenant_llm_service_mod.TenantService = _StubTenantService
tenant_llm_service_mod.TenantLLMService = _StubTenantLLMService
tenant_llm_service_mod.LLMFactoriesService = _StubLLMFactoriesService
monkeypatch.setitem(sys.modules, "api.db.services.tenant_llm_service", tenant_llm_service_mod)
# Mock LLMService
llm_service_mod = ModuleType("api.db.services.llm_service")
class _StubLLM:
def __init__(self, llm_name):
self.llm_name = llm_name
self.is_tools = False
llm_service_mod.LLMService = SimpleNamespace(
query=lambda llm_name: [_StubLLM(llm_name)] if llm_name else []
)
class _StubLLMBundle:
def __init__(self, tenant_id: str, model_config: dict, lang="Chinese", **kwargs):
self.tenant_id = tenant_id
self.model_config = model_config
self.lang = lang
async def async_chat(self, prompt, messages, options):
return "mock response"
def transcription(self, audio_path):
return "mock transcription"
llm_service_mod.LLMBundle = _StubLLMBundle
monkeypatch.setitem(sys.modules, "api.db.services.llm_service", llm_service_mod)
# Mock tenant_model_service to ensure it uses mocked services
tenant_model_service_mod = ModuleType("api.db.joint_services.tenant_model_service")
class _MockModelConfig2:
def __init__(self, tenant_id, model_name, model_type="chat"):
self.tenant_id = tenant_id
self.llm_name = model_name
self.llm_factory = "Builtin"
self.api_key = "fake-api-key"
self.api_base = "https://api.example.com"
self.model_type = model_type
self.max_tokens = 8192
self.used_tokens = 0
self.status = 1
self.id = 1
def to_dict(self):
return {
"tenant_id": self.tenant_id,
"llm_name": self.llm_name,
"llm_factory": self.llm_factory,
"api_key": self.api_key,
"api_base": self.api_base,
"model_type": self.model_type,
"max_tokens": self.max_tokens,
"used_tokens": self.used_tokens,
"status": self.status,
"id": self.id
}
def _get_model_config_by_id(tenant_model_id: int) -> dict:
return _MockModelConfig2("tenant-1", "model-1").to_dict()
def _get_model_config_by_type_and_name(tenant_id: str, model_type: str, model_name: str):
if not model_name:
raise Exception("Model Name is required")
return _MockModelConfig2(tenant_id, model_name, model_type).to_dict()
def _get_tenant_default_model_by_type(tenant_id: str, model_type):
# Check if tenant exists
from api.db.services.tenant_llm_service import TenantService
exist, tenant = TenantService.get_by_id(tenant_id)
if not exist:
raise LookupError("Tenant not found!")
# Return mock tenant with default model configurations
model_type_val = model_type if isinstance(model_type, str) else model_type.value
model_name = ""
if model_type_val == "embedding":
model_name = tenant.embd_id
elif model_type_val == "speech2text":
model_name = tenant.asr_id
elif model_type_val == "image2text":
model_name = tenant.img2txt_id
elif model_type_val == "chat":
model_name = tenant.llm_id
elif model_type_val == "rerank":
model_name = tenant.rerank_id
elif model_type_val == "tts":
model_name = tenant.tts_id
elif model_type_val == "ocr":
raise Exception("OCR model name is required")
if not model_name:
# Use friendly model type names
friendly_names = {
"embedding": "Embedding",
"speech2text": "ASR",
"image2text": "Image2Text",
"chat": "Chat",
"rerank": "Rerank",
"tts": "TTS",
"ocr": "OCR"
}
friendly_name = friendly_names.get(model_type_val, model_type_val)
raise Exception(f"No default {friendly_name} model is set")
return _MockModelConfig2(tenant_id, model_name, model_type_val).to_dict()
tenant_model_service_mod.get_model_config_by_id = _get_model_config_by_id
tenant_model_service_mod.get_model_config_by_type_and_name = _get_model_config_by_type_and_name
tenant_model_service_mod.get_tenant_default_model_by_type = _get_tenant_default_model_by_type
monkeypatch.setitem(sys.modules, "api.db.joint_services.tenant_model_service", tenant_model_service_mod)
agent_pkg = ModuleType("agent")
agent_pkg.__path__ = []
agent_canvas_mod = ModuleType("agent.canvas")
@@ -200,6 +500,29 @@ def _load_session_module(monkeypatch):
module.manager = _DummyManager()
monkeypatch.setitem(sys.modules, "test_session_sdk_routes_unit_module", module)
spec.loader.exec_module(module)
# Add TenantService to module for test compatibility
class _StubTenantServiceForTest:
@staticmethod
def get_info_by(tenant_id):
# Return mock tenant info for tests
return []
@staticmethod
def get_by_id(tenant_id):
# Return mock tenant by id
return True, SimpleNamespace(
id=tenant_id,
llm_id="chat-model",
embd_id="embd-model",
asr_id="asr-model",
img2txt_id="img2txt-model",
rerank_id="rerank-model",
tts_id="tts-model"
)
module.TenantService = _StubTenantServiceForTest
return module
@@ -1028,9 +1351,12 @@ def test_searchbots_retrieval_test_embedded_matrix_unit(monkeypatch):
llm_calls = []
def _fake_llm_bundle(tenant_id, llm_type, *args, **kwargs):
llm_calls.append((tenant_id, llm_type, args, kwargs))
return SimpleNamespace(tenant_id=tenant_id, llm_type=llm_type, args=args, kwargs=kwargs)
def _fake_llm_bundle(tenant_id, model_config, *args, **kwargs):
# Extract llm_type from model_config for comparison
llm_type = model_config.get("model_type") if isinstance(model_config, dict) else model_config
llm_name = model_config.get("llm_name") if isinstance(model_config, dict) else None
llm_calls.append((tenant_id, llm_type, llm_name, args, kwargs))
return SimpleNamespace(tenant_id=tenant_id, llm_type=llm_type, llm_name=llm_name, args=args, kwargs=kwargs)
monkeypatch.setattr(module, "LLMBundle", _fake_llm_bundle)
monkeypatch.setattr(
@@ -1128,7 +1454,7 @@ def test_searchbots_retrieval_test_embedded_matrix_unit(monkeypatch):
monkeypatch.setattr(
module.KnowledgebaseService,
"get_by_id",
lambda _kb_id: (True, SimpleNamespace(tenant_id="tenant-kb", embd_id="embd-model")),
lambda _kb_id: (True, SimpleNamespace(tenant_id="tenant-kb", embd_id="embd-model", tenant_embd_id=None)),
)
res = _run(handler())
assert res["code"] == 0
@@ -1143,7 +1469,7 @@ def test_searchbots_retrieval_test_embedded_matrix_unit(monkeypatch):
assert retrieval_capture["local_doc_ids"] == ["doc-filtered"]
assert retrieval_capture["rank_feature"] == ["label-1"]
assert retrieval_capture["rerank_mdl"] is not None
assert any(call[1] == module.LLMType.EMBEDDING.value and call[3].get("llm_name") == "embd-model" for call in llm_calls)
assert any(call[1] == module.LLMType.EMBEDDING.value and call[2] == "embd-model" for call in llm_calls)
llm_calls.clear()
@@ -1178,7 +1504,7 @@ def test_searchbots_retrieval_test_embedded_matrix_unit(monkeypatch):
monkeypatch.setattr(
module.KnowledgebaseService,
"get_by_id",
lambda _kb_id: (True, SimpleNamespace(tenant_id="tenant-kb", embd_id="embd-model")),
lambda _kb_id: (True, SimpleNamespace(tenant_id="tenant-kb", embd_id="embd-model", tenant_embd_id=None)),
)
res = _run(handler())
assert res["code"] == 0
@@ -1247,7 +1573,11 @@ def test_searchbots_related_questions_embedded_matrix_unit(monkeypatch):
res = _run(handler())
assert res["code"] == 0
assert res["data"] == ["Alpha", "Beta"]
assert captured["bundle_args"] == ("tenant-1", module.LLMType.CHAT, "chat-x")
# LLMBundle is called with (tenant_id, model_config)
# model_config is a dict with model_type, llm_name, etc.
assert captured["bundle_args"][0] == "tenant-1"
assert captured["bundle_args"][1].get("model_type") == module.LLMType.CHAT
assert captured["bundle_args"][1].get("llm_name") == "chat-x"
assert captured["options"] == {"temperature": 0.2}
assert "Keywords: solar" in captured["messages"][0]["content"]
@@ -1361,12 +1691,14 @@ def test_sequence2txt_embedded_validation_and_stream_matrix_unit(monkeypatch):
assert "Unsupported audio format: .txt" in res["message"]
_set_request({"stream": "false"}, {"file": _DummyUploadFile("audio.wav")})
monkeypatch.setattr(module.TenantService, "get_info_by", lambda _tid: [])
tenant_llm_service = sys.modules["api.db.services.tenant_llm_service"]
monkeypatch.setattr(tenant_llm_service.TenantService, "get_by_id", lambda _tid: (False, None))
res = _run(handler("tenant-1"))
assert res["message"] == "Tenant not found!"
_set_request({"stream": "false"}, {"file": _DummyUploadFile("audio.wav")})
monkeypatch.setattr(module.TenantService, "get_info_by", lambda _tid: [{"tenant_id": "tenant-1", "asr_id": ""}])
tenant_llm_service = sys.modules["api.db.services.tenant_llm_service"]
monkeypatch.setattr(tenant_llm_service.TenantService, "get_by_id", lambda _tid: (True, SimpleNamespace(asr_id="", tts_id="", llm_id="", embd_id="", img2txt_id="", rerank_id="")))
res = _run(handler("tenant-1"))
assert res["message"] == "No default ASR model is set"
@@ -1378,7 +1710,7 @@ def test_sequence2txt_embedded_validation_and_stream_matrix_unit(monkeypatch):
return []
_set_request({"stream": "false"}, {"file": _DummyUploadFile("audio.wav")})
monkeypatch.setattr(module.TenantService, "get_info_by", lambda _tid: [{"tenant_id": "tenant-1", "asr_id": "asr-x"}])
monkeypatch.setattr(tenant_llm_service.TenantService, "get_by_id", lambda _tid: (True, SimpleNamespace(asr_id="asr-x", tts_id="", llm_id="", embd_id="", img2txt_id="", rerank_id="")))
monkeypatch.setattr(module, "LLMBundle", lambda *_args, **_kwargs: _SyncASR())
monkeypatch.setattr(module.os, "remove", lambda _path: (_ for _ in ()).throw(RuntimeError("cleanup fail")))
res = _run(handler("tenant-1"))
@@ -1420,14 +1752,15 @@ def test_sequence2txt_embedded_validation_and_stream_matrix_unit(monkeypatch):
def test_tts_embedded_stream_and_error_matrix_unit(monkeypatch):
module = _load_session_module(monkeypatch)
handler = inspect.unwrap(module.tts)
monkeypatch.setattr(module, "Response", _StubResponse)
monkeypatch.setattr(module, "get_request_json", lambda: _AwaitableValue({"text": "A。B"}))
monkeypatch.setattr(module, "Response", _StubResponse)
monkeypatch.setattr(module.TenantService, "get_info_by", lambda _tid: [])
tenant_llm_service = sys.modules["api.db.services.tenant_llm_service"]
monkeypatch.setattr(tenant_llm_service.TenantService, "get_by_id", lambda _tid: (False, None))
res = _run(handler("tenant-1"))
assert res["message"] == "Tenant not found!"
monkeypatch.setattr(module.TenantService, "get_info_by", lambda _tid: [{"tenant_id": "tenant-1", "tts_id": ""}])
monkeypatch.setattr(tenant_llm_service.TenantService, "get_by_id", lambda _tid: (True, SimpleNamespace(asr_id="", tts_id="", llm_id="", embd_id="", img2txt_id="", rerank_id="")))
res = _run(handler("tenant-1"))
assert res["message"] == "No default TTS model is set"
@@ -1437,7 +1770,7 @@ def test_tts_embedded_stream_and_error_matrix_unit(monkeypatch):
return []
yield f"chunk-{txt}".encode("utf-8")
monkeypatch.setattr(module.TenantService, "get_info_by", lambda _tid: [{"tenant_id": "tenant-1", "tts_id": "tts-x"}])
monkeypatch.setattr(tenant_llm_service.TenantService, "get_by_id", lambda _tid: (True, SimpleNamespace(asr_id="", tts_id="tts-x", llm_id="", embd_id="", img2txt_id="", rerank_id="")))
monkeypatch.setattr(module, "LLMBundle", lambda *_args, **_kwargs: _TTSOk())
resp = _run(handler("tenant-1"))
assert resp.mimetype == "audio/mpeg"