diff --git a/src/core/lightrag_based_kb.py b/src/core/lightrag_based_kb.py index bf6750da..9df55877 100644 --- a/src/core/lightrag_based_kb.py +++ b/src/core/lightrag_based_kb.py @@ -120,21 +120,23 @@ class LightRagBasedKB: def _get_llm_func(self, llm_info: dict): """获取 LLM 函数""" - llm_info = llm_info | { - "model_name": "qwen3-1.7b", - "provider": "dashscope" - } - provider_info = config.model_names[llm_info.get("provider")] - api_key = os.getenv(provider_info.get("env")[0] or "OPENAI_API_KEY") or "no_api_key" - base_url = get_docker_safe_url(provider_info.get("base_url", "http://localhost:8081/v1")) + # llm_info = llm_info | { + # "model_name": "qwen3-1.7b", + # "provider": "custom" + # } + # provider_info = config.model_names[llm_info.get("provider")] + # api_key = os.getenv(provider_info.get("env")[0] or "OPENAI_API_KEY") or "no_api_key" + # base_url = get_docker_safe_url(provider_info.get("base_url", "http://localhost:8081/v1")) + from src.models import get_custom_model + llm_info = get_custom_model("qwen3:32b-RFnC") async def llm_model_func(prompt, system_prompt=None, history_messages=[], **kwargs): return await openai_complete_if_cache( - llm_info.get("model_name"), + llm_info.get("name", "qwen3-1.7b"), prompt, system_prompt=system_prompt, history_messages=history_messages, - api_key=api_key, - base_url=base_url, + api_key=llm_info.get("api_key"), + base_url=get_docker_safe_url(llm_info.get("api_base")), extra_body={"enable_thinking": False}, **kwargs, ) @@ -155,42 +157,43 @@ class LightRagBasedKB: ), ) - async def _process_file_to_markdown(self, file_path: str, params=None) -> str: + async def _process_file_to_markdown(self, file_path: str, params: dict | None = None) -> str: """将不同类型的文件转换为 markdown 格式""" - file_path = Path(file_path) - file_ext = file_path.suffix.lower() + file_path_obj = Path(file_path) + file_ext = file_path_obj.suffix.lower() if file_ext == '.pdf': # 使用 OCR 处理 PDF from src.core.indexing import parse_pdf_async - text = await parse_pdf_async(str(file_path), params=params) - return f"Using OCR to process {file_path.name}\n\n{text}" + text = await parse_pdf_async(str(file_path_obj), params=params) + return f"Using OCR to process {file_path_obj.name}\n\n{text}" elif file_ext in ['.txt', '.md']: # 直接读取文本文件 - with open(file_path, 'r', encoding='utf-8') as f: + with open(file_path_obj, 'r', encoding='utf-8') as f: content = f.read() - return f"# {file_path.name}\n\n{content}" + return f"# {file_path_obj.name}\n\n{content}" elif file_ext in ['.doc', '.docx']: # 处理 Word 文档 - from docx import Document - doc = Document(file_path) + + from docx import Document # type: ignore + doc = Document(file_path_obj) text = '\n'.join([para.text for para in doc.paragraphs]) - return f"# {file_path.name}\n\n{text}" + return f"# {file_path_obj.name}\n\n{text}" elif file_ext in ['.jpg', '.jpeg', '.png', '.bmp']: # 使用 OCR 处理图片 - text = ocr.process_image(str(file_path)) - return f"# {file_path.name}\n\n{text}" + text = ocr.process_image(str(file_path_obj)) + return f"# {file_path_obj.name}\n\n{text}" else: # 尝试作为文本文件读取 - import textract - text = textract.process(file_path) - return f"# {file_path.name}\n\n{text}" + import textract # type: ignore + text = textract.process(file_path_obj) + return f"# {file_path_obj.name}\n\n{text}" - async def _process_url_to_markdown(self, url: str, params=None) -> str: + async def _process_url_to_markdown(self, url: str, params: dict | None = None) -> str: """将 URL 转换为 markdown 格式""" import requests from bs4 import BeautifulSoup @@ -231,7 +234,7 @@ class LightRagBasedKB: return {"databases": databases} - def create_database(self, database_name, description, embed_info: dict = None, **kwargs): + def create_database(self, database_name, description, embed_info: dict | None = None, **kwargs): """创建数据库 - data_router.py 使用""" db_id = f"kb_{hashstr(database_name, with_salt=True)}" @@ -285,7 +288,7 @@ class LightRagBasedKB: return {"message": "删除成功"} - async def add_content(self, db_id, items, params=None): + async def add_content(self, db_id, items, params: dict | None = None): """通用的内容添加方法 - 支持文件和URL""" if db_id not in self.databases_meta: raise ValueError(f"Database {db_id} not found") @@ -294,7 +297,7 @@ class LightRagBasedKB: if not rag: raise ValueError(f"Failed to get LightRAG instance for {db_id}") - content_type = params.get('content_type', 'file') + content_type = params.get('content_type', 'file') if params else 'file' processed_items_info = [] @@ -412,6 +415,7 @@ class LightRagBasedKB: if rag: try: # 获取文档的所有 chunks + assert hasattr(rag.text_chunks, 'get_all'), "text_chunks does not have get_all method" all_chunks = await rag.text_chunks.get_all() # 筛选属于该文档的 chunks diff --git a/src/models/__init__.py b/src/models/__init__.py index 885ee059..86b31bfe 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -54,9 +54,7 @@ def select_model(model_provider=None, model_name=None): ) if model_provider == "custom": - model_info = next((x for x in config.custom_models if x["custom_id"] == model_name), None) - if model_info is None: - raise ValueError(f"Model {model_name} not found in custom models") + model_info = get_custom_model(model_name) from src.models.chat_model import CustomModel return CustomModel(model_info) @@ -71,3 +69,12 @@ def select_model(model_provider=None, model_name=None): return model except Exception as e: raise ValueError(f"Model provider {model_provider} load failed, {e} \n {traceback.format_exc()}") + + +def get_custom_model(model_id): + """return model_info""" + assert config.custom_models is not None, "custom_models is not set" + modle_info = next((x for x in config.custom_models if x["custom_id"] == model_id), None) + if modle_info is None: + raise ValueError(f"Model {model_id} not found in custom models") + return modle_info