From 7641ab36997e1594e73c1796fb69eb4fb2179a87 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Mon, 1 Sep 2025 22:37:03 +0800 Subject: [PATCH] =?UTF-8?q?stylle:=20=E4=BD=BF=E7=94=A8=E6=9B=B4=E5=8A=A0?= =?UTF-8?q?=E4=B8=A5=E6=A0=BC=E7=9A=84=E9=A3=8E=E6=A0=BC=E6=A3=80=E6=9F=A5?= =?UTF-8?q?=E4=BB=A5=E5=8F=8A=E4=BB=A3=E7=A0=81=E6=A0=BC=E5=BC=8F=E5=8C=96?= =?UTF-8?q?=E6=96=B9=E6=B3=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Makefile | 13 ++ pyproject.toml | 3 +- scripts/batch_upload.py | 88 +++++--- .../preprocessors/split_data_to_subfiles.py | 21 +- scripts/rename_milvus_collections.py | 14 +- scripts/vllm/test_vllm.py | 3 +- server/db_manager.py | 13 +- server/main.py | 6 +- server/models/kb_models.py | 27 ++- server/models/thread_model.py | 7 +- server/models/user_model.py | 20 +- server/routers/__init__.py | 15 +- server/routers/auth_router.py | 94 ++++----- server/routers/chat_router.py | 144 +++++++------ server/routers/graph_router.py | 159 ++++++--------- server/routers/knowledge_router.py | 154 +++++++------- server/routers/system_router.py | 99 ++++----- server/utils/auth_middleware.py | 23 ++- server/utils/auth_utils.py | 4 +- server/utils/common_utils.py | 25 +-- src/__init__.py | 30 +-- src/agents/__init__.py | 1 + src/agents/chatbot/context.py | 20 +- src/agents/chatbot/graph.py | 29 ++- src/agents/chatbot/state.py | 8 +- src/agents/chatbot/tools.py | 4 +- src/agents/common/base.py | 20 +- src/agents/common/context.py | 42 ++-- src/agents/common/mcp.py | 5 +- src/agents/common/models.py | 11 +- src/agents/common/tools.py | 53 ++--- src/agents/common/utils.py | 20 +- src/agents/react/graph.py | 16 +- src/config/__init__.py | 66 +++--- src/knowledge/chroma_kb.py | 90 ++++----- src/knowledge/graphbase.py | 189 ++++++++++-------- src/knowledge/indexing.py | 90 +++++---- src/knowledge/kb_factory.py | 11 +- src/knowledge/kb_manager.py | 41 ++-- src/knowledge/kb_utils.py | 80 ++++---- src/knowledge/knowledge_base.py | 59 +++--- src/knowledge/lightrag_kb.py | 68 +++---- src/knowledge/milvus_kb.py | 103 ++++------ src/models/__init__.py | 7 +- src/models/chat_model.py | 21 +- src/models/embedding.py | 41 ++-- src/models/rerank_model.py | 26 +-- src/plugins/_ocr.py | 105 ++++------ src/plugins/mineru.py | 94 +++++---- src/plugins/paddlex.py | 106 +++++----- src/utils/__init__.py | 12 +- src/utils/logging_config.py | 16 +- src/utils/prompts.py | 15 +- src/utils/web_search.py | 25 +-- test/test_concurrency.py | 29 +-- test/test_neo4j.py | 5 +- 56 files changed, 1211 insertions(+), 1279 deletions(-) diff --git a/Makefile b/Makefile index 22fb64b4..e7feab7f 100644 --- a/Makefile +++ b/Makefile @@ -12,3 +12,16 @@ logs: @echo "\n\nBranch: $$(git branch --show-current)" @echo "Commit ID: $$(git rev-parse HEAD)" @echo "System: $$(uname -a)" + +###################### +# LINTING AND FORMATTING +###################### + +lint: + uv run python -m ruff check . + uv run python -m ruff format src --diff + uv run python -m ruff check --select I src + +format format_diff: + uv run ruff format + uv run ruff check --select I --fix \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index b7767f32..6798d0d6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,14 +54,13 @@ dependencies = [ "tabulate>=0.9.0", ] [tool.ruff] -line-length = 210 # 代码最大行宽 +line-length = 120 # 代码最大行宽 lint.select = [ # 选择的规则 "F", "E", "W", "UP", ] -lint.ignore = ["F401", "E501"] # 忽略的规则 [dependency-groups] dev = [ diff --git a/scripts/batch_upload.py b/scripts/batch_upload.py index ad908006..680dda40 100644 --- a/scripts/batch_upload.py +++ b/scripts/batch_upload.py @@ -1,4 +1,3 @@ - import asyncio import hashlib import json @@ -7,7 +6,7 @@ import pathlib import httpx import typer from rich.console import Console -from rich.progress import Progress, SpinnerColumn, BarColumn, TextColumn, TimeElapsedColumn +from rich.progress import BarColumn, Progress, SpinnerColumn, TextColumn, TimeElapsedColumn app = typer.Typer() console = Console() @@ -44,12 +43,14 @@ async def upload_file( f"{base_url}/knowledge/files/upload", params={"db_id": db_id}, files=files, - timeout=300, # 5 minutes timeout for large files + timeout=300, # 5 minutes timeout for large files ) response.raise_for_status() return response.json().get("file_path") except httpx.HTTPStatusError as e: - console.print(f"[bold red]Failed to upload {file_path.name}: {e.response.status_code} - {e.response.text}[/bold red]") + console.print( + f"[bold red]Failed to upload {file_path.name}: {e.response.status_code} - {e.response.text}[/bold red]" + ) return None except httpx.RequestError as e: console.print(f"[bold red]Failed to upload {file_path.name}: {e}[/bold red]") @@ -75,21 +76,23 @@ async def process_document( "enable_ocr": enable_ocr, "use_qa_split": use_qa_split, "qa_separator": qa_separator, - "content_type": "file" + "content_type": "file", } try: response = await client.post( f"{base_url}/knowledge/databases/{db_id}/documents", json={"items": [server_file_path], "params": params}, - timeout=600, # 10 minutes timeout for processing + timeout=600, # 10 minutes timeout for processing ) response.raise_for_status() result = response.json() # Check if the overall request was successful if result.get("status") != "success": - console.print(f"[bold yellow]Processing warning for {server_file_path}: {result.get('message')}[/bold yellow]") + console.print( + f"[bold yellow]Processing warning for {server_file_path}: {result.get('message')}[/bold yellow]" + ) return False # Check the specific file's processing status in the items array @@ -129,7 +132,9 @@ async def process_document( return False except httpx.HTTPStatusError as e: - console.print(f"[bold red]Failed to process {server_file_path}: {e.response.status_code} - {e.response.text}[/bold red]") + console.print( + f"[bold red]Failed to process {server_file_path}: {e.response.status_code} - {e.response.text}[/bold red]" + ) return False except httpx.RequestError as e: console.print(f"[bold red]Failed to process {server_file_path}: {e}[/bold red]") @@ -159,17 +164,20 @@ async def worker( progress.update(upload_task_id, advance=1, postfix=f"Uploaded {file_path.name}") if not server_file_path: - progress.update(process_task_id, advance=1) # Mark as processed to not hang the progress bar + progress.update(process_task_id, advance=1) # Mark as processed to not hang the progress bar return file_path, file_hash, "upload_failed" # 2. Process file success = await process_document( - client, base_url, db_id, server_file_path, + client, + base_url, + db_id, + server_file_path, enable_ocr=enable_ocr, chunk_size=chunk_size, chunk_overlap=chunk_overlap, use_qa_split=use_qa_split, - qa_separator=qa_separator + qa_separator=qa_separator, ) progress.update(process_task_id, advance=1, postfix=f"Processed {file_path.name}") @@ -193,7 +201,7 @@ def load_processed_files(record_file: pathlib.Path) -> set[str]: try: with open(record_file) as f: data = json.load(f) - return set(data.get('processed_files', [])) + return set(data.get("processed_files", [])) except (OSError, json.JSONDecodeError) as e: console.print(f"[bold yellow]Warning: Could not load processed files record: {e}[/bold yellow]") return set() @@ -205,8 +213,8 @@ def save_processed_files(record_file: pathlib.Path, processed_files: set[str]): record_file.parent.mkdir(parents=True, exist_ok=True) try: - with open(record_file, 'w') as f: - json.dump({'processed_files': list(processed_files)}, f, indent=2) + with open(record_file, "w") as f: + json.dump({"processed_files": list(processed_files)}, f, indent=2) except OSError as e: console.print(f"[bold red]Error: Could not save processed files record: {e}[/bold red]") @@ -232,7 +240,9 @@ async def convert_to_markdown( console.print(f"[bold red]Failed to convert {server_file_path}: {result.get('message')}[/bold red]") return None except httpx.HTTPStatusError as e: - console.print(f"[bold red]Failed to convert {server_file_path}: {e.response.status_code} - {e.response.text}[/bold red]") + console.print( + f"[bold red]Failed to convert {server_file_path}: {e.response.status_code} - {e.response.text}[/bold red]" + ) return None except httpx.RequestError as e: console.print(f"[bold red]Request failed for {server_file_path}: {e}[/bold red]") @@ -267,7 +277,7 @@ async def trans_worker( try: output_path = output_dir / file_path.with_suffix(".md").name output_path.parent.mkdir(parents=True, exist_ok=True) - with open(output_path, 'w', encoding='utf-8') as f: + with open(output_path, "w", encoding="utf-8") as f: f.write(markdown_content) progress.update(task_id, advance=1, postfix=f"[green]Converted {file_path.name}[/green]") return file_path, "success" @@ -280,14 +290,18 @@ async def trans_worker( @app.command() def upload( db_id: str = typer.Option(..., help="The ID of the knowledge base."), - directory: pathlib.Path = typer.Option(..., help="The directory containing files to upload.", exists=True, file_okay=False), + directory: pathlib.Path = typer.Option( + ..., help="The directory containing files to upload.", exists=True, file_okay=False + ), pattern: str = typer.Option("*.md", help="The glob pattern for files to upload (e.g., '*.pdf', '**/*.txt')."), base_url: str = typer.Option("http://127.0.0.1:5050/api", help="The base URL of the API server."), username: str = typer.Option(..., help="Admin username for login."), password: str = typer.Option(..., help="Admin password for login."), concurrency: int = typer.Option(1, help="The number of concurrent upload/process tasks."), recursive: bool = typer.Option(False, "--recursive", "-r", help="Search for files recursively in subdirectories."), - record_file: pathlib.Path = typer.Option("scripts/tmp/batch_processed_files.txt", help="File to store processed files record."), + record_file: pathlib.Path = typer.Option( + "scripts/tmp/batch_processed_files.txt", help="File to store processed files record." + ), chunk_size: int = typer.Option(1000, help="Chunk size for document processing."), chunk_overlap: int = typer.Option(200, help="Chunk overlap for document processing."), enable_ocr: str = typer.Option("paddlex_ocr", help="OCR engine to use (paddlex_ocr, mineru_ocr, disable)."), @@ -325,7 +339,9 @@ def upload( files_to_upload.append((file_path, file_hash)) if not files_to_upload: - console.print(f"[bold green]All {len(all_files)} files have already been processed. Nothing to do.[/bold green]") + console.print( + f"[bold green]All {len(all_files)} files have already been processed. Nothing to do.[/bold green]" + ) raise typer.Exit() console.print(f"Found {len(all_files)} total files:") @@ -361,13 +377,20 @@ def upload( for file_path, file_hash in files_to_upload: task = asyncio.create_task( worker( - semaphore, client, base_url, db_id, file_path, file_hash, - progress, upload_task_id, process_task_id, + semaphore, + client, + base_url, + db_id, + file_path, + file_hash, + progress, + upload_task_id, + process_task_id, enable_ocr=enable_ocr, chunk_size=chunk_size, chunk_overlap=chunk_overlap, use_qa_split=use_qa_split, - qa_separator=qa_separator + qa_separator=qa_separator, ) ) tasks.append(task) @@ -381,12 +404,12 @@ def upload( newly_processed_hashes = set() for file_path, file_hash, status in results: - if status == 'success': + if status == "success": successful_files.append(file_path) newly_processed_hashes.add(file_hash) - elif status == 'upload_failed': + elif status == "upload_failed": upload_failures.append(file_path) - elif status == 'processing_failed': + elif status == "processing_failed": processing_failures.append(file_path) # Don't add to processed files if processing failed @@ -394,7 +417,9 @@ def upload( if newly_processed_hashes: all_processed_files = processed_files | newly_processed_hashes save_processed_files(record_file, all_processed_files) - console.print(f"[bold green]Updated processed files record with {len(newly_processed_hashes)} new entries.[/bold green]") + console.print( + f"[bold green]Updated processed files record with {len(newly_processed_hashes)} new entries.[/bold green]" + ) console.print("[bold green]Batch operation complete.[/bold green]") console.print(f" - [green]Successful:[/green] {len(successful_files)}") @@ -413,7 +438,9 @@ def upload( @app.command() def trans( db_id: str = typer.Option(..., help="The ID of the knowledge base (for temporary file upload)."), - directory: pathlib.Path = typer.Option(..., help="The directory containing files to convert.", exists=True, file_okay=False), + directory: pathlib.Path = typer.Option( + ..., help="The directory containing files to convert.", exists=True, file_okay=False + ), output_dir: pathlib.Path = typer.Option("output_markdown", help="The directory to save converted markdown files."), pattern: str = typer.Option("*.docx", help="The glob pattern for files to convert (e.g., '*.pdf', '*.docx')."), base_url: str = typer.Option("http://127.0.0.1:5050/api", help="The base URL of the API server."), @@ -467,10 +494,7 @@ def trans( for file_path in files_to_convert: task = asyncio.create_task( - trans_worker( - semaphore, client, base_url, db_id, file_path, - output_dir, progress, task_id - ) + trans_worker(semaphore, client, base_url, db_id, file_path, output_dir, progress, task_id) ) tasks.append(task) @@ -481,7 +505,7 @@ def trans( failed_files = [] for file_path, status in results: - if status == 'success': + if status == "success": successful_files.append(file_path) else: failed_files.append((file_path, status)) diff --git a/scripts/preprocessors/split_data_to_subfiles.py b/scripts/preprocessors/split_data_to_subfiles.py index f74a1072..6e132f3d 100644 --- a/scripts/preprocessors/split_data_to_subfiles.py +++ b/scripts/preprocessors/split_data_to_subfiles.py @@ -1,26 +1,30 @@ -import typer -import pandas as pd import json import random import re from pathlib import Path +import pandas as pd +import typer + app = typer.Typer() + def sanitize_filename(name: str) -> str: return re.sub(r'[\\/*?:"<>|]', "_", str(name).strip()) + def random_suffix() -> str: return f"_{random.randint(10000000, 99999999)}" + def read_table(file_path: Path) -> pd.DataFrame: suffix = file_path.suffix.lower() - if suffix in ['.xlsx', '.xls']: + if suffix in [".xlsx", ".xls"]: return pd.read_excel(file_path) - elif suffix == '.csv': + elif suffix == ".csv": return pd.read_csv(file_path) - elif suffix == '.json': - with open(file_path, encoding='utf-8') as f: + elif suffix == ".json": + with open(file_path, encoding="utf-8") as f: data = json.load(f) if isinstance(data, list) and len(data) > 1 and isinstance(data[0], dict): return pd.DataFrame(data) @@ -29,6 +33,7 @@ def read_table(file_path: Path) -> pd.DataFrame: else: raise ValueError(f"不支持的文件格式:{suffix}") + def export_txts(df: pd.DataFrame, output_dir: Path, title_field: str = "标题"): output_dir.mkdir(parents=True, exist_ok=True) df.columns = [c.strip() for c in df.columns] @@ -54,11 +59,12 @@ def export_txts(df: pd.DataFrame, output_dir: Path, title_field: str = "标题") typer.echo(f"✅ 成功导出 {len(df)} 个文件到目录:{output_dir}") + @app.command() def convert( input_file: Path = typer.Argument(..., help="输入文件(.xlsx/.xls/.csv/.json)"), out_dir: Path = typer.Option("output", help="输出目录"), - title_field: str = typer.Option("标题", help="标题字段名(用于文件名)") + title_field: str = typer.Option("标题", help="标题字段名(用于文件名)"), ): """ 将结构化数据文件(Excel/CSV/JSON)转换为多个 .txt 文件。 @@ -70,5 +76,6 @@ def convert( typer.echo(f"❌ 错误:{e}", err=True) raise typer.Exit(code=1) + if __name__ == "__main__": app() diff --git a/scripts/rename_milvus_collections.py b/scripts/rename_milvus_collections.py index 1baa5eaa..8f23f284 100644 --- a/scripts/rename_milvus_collections.py +++ b/scripts/rename_milvus_collections.py @@ -1,5 +1,7 @@ import os -from pymilvus import utility, connections, Collection + +from pymilvus import Collection, connections, utility + def get_collection_info(collection_name, alias): """Safely gets a collection object and its number of entities.""" @@ -11,13 +13,14 @@ def get_collection_info(collection_name, alias): print(f"Error getting info for collection '{collection_name}': {e}") return None, 0 + def rename_and_resolve_duplicates(): """ Connects to Milvus, renames collections from 'kb_kb_' to 'kb_', and resolves duplicates by keeping the collection with more rows. """ - milvus_uri = os.getenv('MILVUS_URI', 'http://localhost:19530') - milvus_token = os.getenv('MILVUS_TOKEN', '') + milvus_uri = os.getenv("MILVUS_URI", "http://localhost:19530") + milvus_token = os.getenv("MILVUS_TOKEN", "") connection_alias = "rename_script" try: @@ -26,7 +29,7 @@ def rename_and_resolve_duplicates(): print("Successfully connected to Milvus.") all_collections = utility.list_collections(using=connection_alias) - collections_to_rename = [c for c in all_collections if c.startswith('kb_kb_')] + collections_to_rename = [c for c in all_collections if c.startswith("kb_kb_")] if not collections_to_rename: print("No collections with the prefix 'kb_kb_' found. Nothing to do.") @@ -35,7 +38,7 @@ def rename_and_resolve_duplicates(): print(f"Found {len(collections_to_rename)} collections with 'kb_kb_' prefix to process.") for old_name in collections_to_rename: - new_name = old_name.replace('kb_kb_', 'kb_', 1) + new_name = old_name.replace("kb_kb_", "kb_", 1) try: print(f"Attempting to rename '{old_name}' to '{new_name}'...") utility.rename_collection(old_name, new_name, using=connection_alias) @@ -74,5 +77,6 @@ def rename_and_resolve_duplicates(): connections.disconnect(connection_alias) print("Disconnected from Milvus.") + if __name__ == "__main__": rename_and_resolve_duplicates() diff --git a/scripts/vllm/test_vllm.py b/scripts/vllm/test_vllm.py index 204a4a76..2b0b92c7 100644 --- a/scripts/vllm/test_vllm.py +++ b/scripts/vllm/test_vllm.py @@ -1,4 +1,5 @@ from openai import OpenAI + # Set OpenAI's API key and API base to use vLLM's API server. openai_api_key = "EMPTY" openai_api_base = "http://localhost:8080/v1" @@ -13,6 +14,6 @@ chat_response = client.chat.completions.create( messages=[ {"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "Tell me a joke."}, - ] + ], ) print("Chat response:", chat_response) diff --git a/server/db_manager.py b/server/db_manager.py index 24308c71..b40056a6 100644 --- a/server/db_manager.py +++ b/server/db_manager.py @@ -1,16 +1,18 @@ import os import pathlib -from sqlalchemy import create_engine -from sqlalchemy.orm import sessionmaker from contextlib import contextmanager -from src import config +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker + from server.models import Base -from server.models.user_model import User -from server.models.thread_model import Thread from server.models.kb_models import KnowledgeDatabase, KnowledgeFile, KnowledgeNode +from server.models.thread_model import Thread +from server.models.user_model import User +from src import config from src.utils import logger + class DBManager: """数据库管理器 - 只提供基础的数据库连接和会话管理""" @@ -65,5 +67,6 @@ class DBManager: finally: session.close() + # 创建全局数据库管理器实例 db_manager = DBManager() diff --git a/server/main.py b/server/main.py index d6e76bac..aeb6c31d 100644 --- a/server/main.py +++ b/server/main.py @@ -1,6 +1,5 @@ import uvicorn - -from fastapi import FastAPI, Request, HTTPException, status, Depends +from fastapi import Depends, FastAPI, HTTPException, Request, status from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse from starlette.middleware.base import BaseHTTPMiddleware @@ -24,6 +23,7 @@ app.add_middleware( allow_headers=["*"], ) + # 鉴权中间件 class AuthMiddleware(BaseHTTPMiddleware): async def dispatch(self, request: Request, call_next): @@ -56,9 +56,9 @@ class AuthMiddleware(BaseHTTPMiddleware): # 继续处理请求 return await call_next(request) + # 添加鉴权中间件 app.add_middleware(AuthMiddleware) if __name__ == "__main__": uvicorn.run(app, host="0.0.0.0", port=5050, threads=10, workers=10, reload=True) - diff --git a/server/models/kb_models.py b/server/models/kb_models.py index c877ec17..65617a21 100644 --- a/server/models/kb_models.py +++ b/server/models/kb_models.py @@ -1,13 +1,16 @@ -from sqlalchemy import Column, Integer, String, DateTime, JSON, ForeignKey, Text +import time + +from sqlalchemy import JSON, Column, DateTime, ForeignKey, Integer, String, Text from sqlalchemy.orm import relationship from sqlalchemy.sql import func -import time from server.models import Base + class KnowledgeDatabase(Base): """知识库模型""" - __tablename__ = 'knowledge_databases' + + __tablename__ = "knowledge_databases" id = Column(Integer, primary_key=True, autoincrement=True) db_id = Column(String, nullable=False, unique=True, index=True) # 数据库ID @@ -31,7 +34,7 @@ class KnowledgeDatabase(Base): "embed_model": self.embed_model, "dimension": self.dimension, "metadata": self.meta_info or {}, - "created_at": self.created_at.isoformat() if self.created_at else None + "created_at": self.created_at.isoformat() if self.created_at else None, } # 添加文件信息 if self.files: @@ -40,13 +43,15 @@ class KnowledgeDatabase(Base): result["files"] = {} return result + class KnowledgeFile(Base): """知识库文件模型""" - __tablename__ = 'knowledge_files' + + __tablename__ = "knowledge_files" id = Column(Integer, primary_key=True, autoincrement=True) file_id = Column(String, nullable=False, index=True) # 文件ID - database_id = Column(String, ForeignKey('knowledge_databases.db_id'), nullable=False) # 所属数据库ID + database_id = Column(String, ForeignKey("knowledge_databases.db_id"), nullable=False) # 所属数据库ID filename = Column(String, nullable=False) # 文件名 path = Column(String, nullable=False) # 文件路径 file_type = Column(String, nullable=False) # 文件类型 @@ -71,18 +76,20 @@ class KnowledgeFile(Base): "type": self.file_type, "status": self.status, "node_count": self.computed_node_count, - "created_at": self.created_at.timestamp() if self.created_at else time.time() + "created_at": self.created_at.timestamp() if self.created_at else time.time(), } if with_nodes: result["nodes"] = [node.to_dict() for node in self.nodes] if self.nodes else [] return result + class KnowledgeNode(Base): """知识块模型""" - __tablename__ = 'knowledge_nodes' + + __tablename__ = "knowledge_nodes" id = Column(Integer, primary_key=True, autoincrement=True) - file_id = Column(String, ForeignKey('knowledge_files.file_id'), nullable=False) # 所属文件ID + file_id = Column(String, ForeignKey("knowledge_files.file_id"), nullable=False) # 所属文件ID text = Column(Text, nullable=False) # 文本内容 hash = Column(String, nullable=True) # 文本哈希值 start_char_idx = Column(Integer, nullable=True) # 开始字符索引 @@ -101,5 +108,5 @@ class KnowledgeNode(Base): "hash": self.hash, "start_char_idx": self.start_char_idx, "end_char_idx": self.end_char_idx, - "metadata": self.meta_info or {} # 确保映射正确 + "metadata": self.meta_info or {}, # 确保映射正确 } diff --git a/server/models/thread_model.py b/server/models/thread_model.py index 10008373..4bdf764b 100644 --- a/server/models/thread_model.py +++ b/server/models/thread_model.py @@ -1,13 +1,14 @@ -from sqlalchemy import Column, String, Integer, DateTime, Text, ForeignKey -from sqlalchemy.sql import func -from sqlalchemy.orm import relationship +from sqlalchemy import Column, DateTime, ForeignKey, Integer, String, Text from sqlalchemy.dialects.mysql import JSON +from sqlalchemy.orm import relationship +from sqlalchemy.sql import func from server.models import Base class Thread(Base): """对话线程表""" + __tablename__ = "thread" id = Column(String(64), primary_key=True, index=True, comment="线程ID") diff --git a/server/models/user_model.py b/server/models/user_model.py index 32a5ae88..fc937787 100644 --- a/server/models/user_model.py +++ b/server/models/user_model.py @@ -1,17 +1,19 @@ -from sqlalchemy import Column, Integer, String, DateTime, ForeignKey, Text -from sqlalchemy.sql import func +from sqlalchemy import Column, DateTime, ForeignKey, Integer, String, Text from sqlalchemy.orm import relationship +from sqlalchemy.sql import func from server.models import Base + class User(Base): """用户模型""" - __tablename__ = 'users' + + __tablename__ = "users" id = Column(Integer, primary_key=True, autoincrement=True) username = Column(String, nullable=False, unique=True, index=True) password_hash = Column(String, nullable=False) - role = Column(String, nullable=False, default='user') # 角色: superadmin, admin, user + role = Column(String, nullable=False, default="user") # 角色: superadmin, admin, user created_at = Column(DateTime, default=func.now()) last_login = Column(DateTime, nullable=True) @@ -24,18 +26,20 @@ class User(Base): "username": self.username, "role": self.role, "created_at": self.created_at.isoformat() if self.created_at else None, - "last_login": self.last_login.isoformat() if self.last_login else None + "last_login": self.last_login.isoformat() if self.last_login else None, } if include_password: result["password_hash"] = self.password_hash return result + class OperationLog(Base): """操作日志模型""" - __tablename__ = 'operation_logs' + + __tablename__ = "operation_logs" id = Column(Integer, primary_key=True, autoincrement=True) - user_id = Column(Integer, ForeignKey('users.id'), nullable=False) + user_id = Column(Integer, ForeignKey("users.id"), nullable=False) operation = Column(String, nullable=False) details = Column(Text, nullable=True) ip_address = Column(String, nullable=True) @@ -51,5 +55,5 @@ class OperationLog(Base): "operation": self.operation, "details": self.details, "ip_address": self.ip_address, - "timestamp": self.timestamp.isoformat() if self.timestamp else None + "timestamp": self.timestamp.isoformat() if self.timestamp else None, } diff --git a/server/routers/__init__.py b/server/routers/__init__.py index 48dd6cd6..bfc7b8b2 100644 --- a/server/routers/__init__.py +++ b/server/routers/__init__.py @@ -1,15 +1,16 @@ from fastapi import APIRouter -from server.routers.system_router import system + from server.routers.auth_router import auth from server.routers.chat_router import chat -from server.routers.knowledge_router import knowledge from server.routers.graph_router import graph +from server.routers.knowledge_router import knowledge +from server.routers.system_router import system router = APIRouter() # 注册路由结构 -router.include_router(system) # /api/system/* -router.include_router(auth) # /api/auth/* -router.include_router(chat) # /api/chat/* -router.include_router(knowledge) # /api/knowledge/* -router.include_router(graph) # /api/graph/* +router.include_router(system) # /api/system/* +router.include_router(auth) # /api/auth/* +router.include_router(chat) # /api/chat/* +router.include_router(knowledge) # /api/knowledge/* +router.include_router(graph) # /api/graph/* diff --git a/server/routers/auth_router.py b/server/routers/auth_router.py index 51689741..5135ea18 100644 --- a/server/routers/auth_router.py +++ b/server/routers/auth_router.py @@ -1,18 +1,20 @@ +from datetime import datetime + from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi.security import OAuth2PasswordRequestForm from pydantic import BaseModel from sqlalchemy.orm import Session -from datetime import datetime from server.db_manager import db_manager -from server.models.user_model import User, OperationLog +from server.models.user_model import OperationLog, User +from server.utils.auth_middleware import get_admin_user, get_current_user, get_db, get_superadmin_user, oauth2_scheme from server.utils.auth_utils import AuthUtils -from server.utils.auth_middleware import get_db, get_current_user, get_admin_user, get_superadmin_user, oauth2_scheme from server.utils.common_utils import log_operation # 创建路由器 auth = APIRouter(prefix="/auth", tags=["authentication"]) + # 请求和响应模型 class Token(BaseModel): access_token: str @@ -21,16 +23,19 @@ class Token(BaseModel): username: str role: str + class UserCreate(BaseModel): username: str password: str role: str = "user" + class UserUpdate(BaseModel): username: str | None = None password: str | None = None role: str | None = None + class UserResponse(BaseModel): id: int username: str @@ -38,10 +43,12 @@ class UserResponse(BaseModel): created_at: str last_login: str | None = None + class InitializeAdmin(BaseModel): username: str password: str + # ============================================================================= # === 工具函数 === # ============================================================================= @@ -52,11 +59,9 @@ class InitializeAdmin(BaseModel): # === 认证分组 === # ============================================================================= + @auth.post("/token", response_model=Token) -async def login_for_access_token( - form_data: OAuth2PasswordRequestForm = Depends(), - db: Session = Depends(get_db) -): +async def login_for_access_token(form_data: OAuth2PasswordRequestForm = Depends(), db: Session = Depends(get_db)): # 查找用户 user = db.query(User).filter(User.username == form_data.username).first() @@ -84,21 +89,20 @@ async def login_for_access_token( "token_type": "bearer", "user_id": user.id, "username": user.username, - "role": user.role + "role": user.role, } + # 路由:校验是否需要初始化管理员 @auth.get("/check-first-run") async def check_first_run(): is_first_run = db_manager.check_first_run() return {"first_run": is_first_run} + # 路由:初始化管理员账户 @auth.post("/initialize", response_model=Token) -async def initialize_admin( - admin_data: InitializeAdmin, - db: Session = Depends(get_db) -): +async def initialize_admin(admin_data: InitializeAdmin, db: Session = Depends(get_db)): # 检查是否是首次运行 if not db_manager.check_first_run(): raise HTTPException( @@ -110,10 +114,7 @@ async def initialize_admin( hashed_password = AuthUtils.hash_password(admin_data.password) new_admin = User( - username=admin_data.username, - password_hash=hashed_password, - role="superadmin", - last_login=datetime.now() + username=admin_data.username, password_hash=hashed_password, role="superadmin", last_login=datetime.now() ) db.add(new_admin) @@ -132,29 +133,30 @@ async def initialize_admin( "token_type": "bearer", "user_id": new_admin.id, "username": new_admin.username, - "role": new_admin.role + "role": new_admin.role, } + # 路由:获取当前用户信息 # ============================================================================= # === 用户信息分组 === # ============================================================================= + @auth.get("/me", response_model=UserResponse) async def read_users_me(current_user: User = Depends(get_current_user)): return current_user.to_dict() + # 路由:创建新用户(管理员权限) # ============================================================================= # === 用户管理分组 === # ============================================================================= + @auth.post("/users", response_model=UserResponse) async def create_user( - user_data: UserCreate, - request: Request, - current_user: User = Depends(get_admin_user), - db: Session = Depends(get_db) + user_data: UserCreate, request: Request, current_user: User = Depends(get_admin_user), db: Session = Depends(get_db) ): # 检查用户名是否已存在 existing_user = db.query(User).filter(User.username == user_data.username).first() @@ -182,45 +184,30 @@ async def create_user( detail="管理员只能创建普通用户账户", ) - new_user = User( - username=user_data.username, - password_hash=hashed_password, - role=user_data.role - ) + new_user = User(username=user_data.username, password_hash=hashed_password, role=user_data.role) db.add(new_user) db.commit() db.refresh(new_user) # 记录操作 - log_operation( - db, - current_user.id, - "创建用户", - f"创建用户: {user_data.username}, 角色: {user_data.role}", - request - ) + log_operation(db, current_user.id, "创建用户", f"创建用户: {user_data.username}, 角色: {user_data.role}", request) return new_user.to_dict() + # 路由:获取所有用户(管理员权限) @auth.get("/users", response_model=list[UserResponse]) async def read_users( - skip: int = 0, - limit: int = 100, - current_user: User = Depends(get_admin_user), - db: Session = Depends(get_db) + skip: int = 0, limit: int = 100, current_user: User = Depends(get_admin_user), db: Session = Depends(get_db) ): users = db.query(User).offset(skip).limit(limit).all() return [user.to_dict() for user in users] + # 路由:获取特定用户信息(管理员权限) @auth.get("/users/{user_id}", response_model=UserResponse) -async def read_user( - user_id: int, - current_user: User = Depends(get_admin_user), - db: Session = Depends(get_db) -): +async def read_user(user_id: int, current_user: User = Depends(get_admin_user), db: Session = Depends(get_db)): user = db.query(User).filter(User.id == user_id).first() if user is None: raise HTTPException( @@ -229,6 +216,7 @@ async def read_user( ) return user.to_dict() + # 路由:更新用户信息(管理员权限) @auth.put("/users/{user_id}", response_model=UserResponse) async def update_user( @@ -236,7 +224,7 @@ async def update_user( user_data: UserUpdate, request: Request, current_user: User = Depends(get_admin_user), - db: Session = Depends(get_db) + db: Session = Depends(get_db), ): user = db.query(User).filter(User.id == user_id).first() if user is None: @@ -284,23 +272,15 @@ async def update_user( db.commit() # 记录操作 - log_operation( - db, - current_user.id, - "更新用户", - f"更新用户ID {user_id}: {', '.join(update_details)}", - request - ) + log_operation(db, current_user.id, "更新用户", f"更新用户ID {user_id}: {', '.join(update_details)}", request) return user.to_dict() + # 路由:删除用户(管理员权限) @auth.delete("/users/{user_id}", response_model=dict) async def delete_user( - user_id: int, - request: Request, - current_user: User = Depends(get_admin_user), - db: Session = Depends(get_db) + user_id: int, request: Request, current_user: User = Depends(get_admin_user), db: Session = Depends(get_db) ): user = db.query(User).filter(User.id == user_id).first() if user is None: @@ -335,11 +315,7 @@ async def delete_user( # 记录操作 log_operation( - db, - current_user.id, - "删除用户", - f"删除用户: {user.username}, ID: {user.id}, 角色: {user.role}", - request + db, current_user.id, "删除用户", f"删除用户: {user.username}, ID: {user.id}, 角色: {user.role}", request ) # 删除用户 diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index 45f58e32..9e548476 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -1,24 +1,25 @@ -import os -import json import asyncio +import json +import os +import time import traceback import uuid -import time + from fastapi import APIRouter, Body, Depends, HTTPException, Query from fastapi.responses import StreamingResponse from langchain_core.messages import AIMessageChunk, HumanMessage -from sqlalchemy.orm import Session from pydantic import BaseModel +from sqlalchemy.orm import Session -from src import executor, config +from server.models.thread_model import Thread +from server.models.user_model import User +from server.routers.auth_router import get_admin_user +from server.utils.auth_middleware import get_db, get_required_user +from src import config, executor from src.agents import agent_manager +from src.agents.common.tools import gen_tool_info, get_buildin_tools from src.models import select_model from src.utils.logging_config import logger -from src.agents.common.tools import get_buildin_tools, gen_tool_info -from server.routers.auth_router import get_admin_user -from server.utils.auth_middleware import get_required_user, get_db -from server.models.user_model import User -from server.models.thread_model import Thread chat = APIRouter(prefix="/chat", tags=["chat"]) @@ -26,6 +27,7 @@ chat = APIRouter(prefix="/chat", tags=["chat"]) # > === 智能体管理分组 === # ============================================================================= + @chat.get("/default_agent") async def get_default_agent(current_user: User = Depends(get_required_user)): """获取默认智能体ID(需要登录)""" @@ -42,8 +44,9 @@ async def get_default_agent(current_user: User = Depends(get_required_user)): logger.error(f"获取默认智能体出错: {e}") raise HTTPException(status_code=500, detail=f"获取默认智能体出错: {str(e)}") + @chat.post("/set_default_agent") -async def set_default_agent(request_data: dict = Body(...), current_user = Depends(get_admin_user)): +async def set_default_agent(request_data: dict = Body(...), current_user=Depends(get_admin_user)): """设置默认智能体ID (仅管理员)""" try: agent_id = request_data.get("agent_id") @@ -69,15 +72,18 @@ async def set_default_agent(request_data: dict = Body(...), current_user = Depen logger.error(f"设置默认智能体出错: {e}") raise HTTPException(status_code=500, detail=f"设置默认智能体出错: {str(e)}") + # ============================================================================= # > === 对话分组 === # ============================================================================= + @chat.post("/call") async def call(query: str = Body(...), meta: dict = Body(None), current_user: User = Depends(get_required_user)): """调用模型进行简单问答(需要登录)""" meta = meta or {} model = select_model(model_provider=meta.get("model_provider"), model_name=meta.get("model_name")) + async def predict_async(query): loop = asyncio.get_event_loop() return await loop.run_in_executor(executor, model.predict, query) @@ -87,6 +93,7 @@ async def call(query: str = Body(...), meta: dict = Body(None), current_user: Us return {"response": response.content} + @chat.get("/agent") async def get_agent(current_user: User = Depends(get_required_user)): """获取所有可用智能体(需要登录)""" @@ -94,35 +101,39 @@ async def get_agent(current_user: User = Depends(get_required_user)): # logger.debug(f"agents: {agents}") return {"agents": agents} + @chat.post("/agent/{agent_id}") -async def chat_agent(agent_id: str, - query: str = Body(...), - config: dict = Body({}), - meta: dict = Body({}), - current_user: User = Depends(get_required_user)): +async def chat_agent( + agent_id: str, + query: str = Body(...), + config: dict = Body({}), + meta: dict = Body({}), + current_user: User = Depends(get_required_user), +): """使用特定智能体进行对话(需要登录)""" logger.info(f"agent_id: {agent_id}, query: {query}, config: {config}, meta: {meta}") - meta.update({ - "query": query, - "agent_id": agent_id, - "server_model_name": config.get("model", agent_id), - "thread_id": config.get("thread_id"), - "user_id": current_user.id - }) + meta.update( + { + "query": query, + "agent_id": agent_id, + "server_model_name": config.get("model", agent_id), + "thread_id": config.get("thread_id"), + "user_id": current_user.id, + } + ) # 将meta和thread_id整合到config中 def make_chunk(content=None, **kwargs): - - return json.dumps({ - "request_id": meta.get("request_id"), - "response": content, - **kwargs - }, ensure_ascii=False).encode('utf-8') + b"\n" + return ( + json.dumps( + {"request_id": meta.get("request_id"), "response": content, **kwargs}, ensure_ascii=False + ).encode("utf-8") + + b"\n" + ) async def stream_messages(): - # 代表服务端已经收到了请求 yield make_chunk(status="init", meta=meta, msg=HumanMessage(content=query).model_dump()) @@ -145,39 +156,38 @@ async def chat_agent(agent_id: str, async for msg, metadata in agent.stream_messages(messages, input_context=input_context): # logger.debug(f"msg: {msg.model_dump()}, metadata: {metadata}") if isinstance(msg, AIMessageChunk): - yield make_chunk(content=msg.content, - msg=msg.model_dump(), - metadata=metadata, - status="loading") + yield make_chunk(content=msg.content, msg=msg.model_dump(), metadata=metadata, status="loading") else: - yield make_chunk(msg=msg.model_dump(), - metadata=metadata, - status="loading") + yield make_chunk(msg=msg.model_dump(), metadata=metadata, status="loading") yield make_chunk(status="finished", meta=meta) except Exception as e: logger.error(f"Error streaming messages: {e}, {traceback.format_exc()}") yield make_chunk(message=f"Error streaming messages: {e}", status="error") - return StreamingResponse(stream_messages(), media_type='application/json') + return StreamingResponse(stream_messages(), media_type="application/json") + # ============================================================================= # > === 模型管理分组 === # ============================================================================= + @chat.get("/models") async def get_chat_models(model_provider: str, current_user: User = Depends(get_admin_user)): """获取指定模型提供商的模型列表(需要登录)""" model = select_model(model_provider=model_provider) return {"models": model.get_models()} + @chat.post("/models/update") -async def update_chat_models(model_provider: str, model_names: list[str], current_user = Depends(get_admin_user)): +async def update_chat_models(model_provider: str, model_names: list[str], current_user=Depends(get_admin_user)): """更新指定模型提供商的模型列表 (仅管理员)""" config.model_names[model_provider]["models"] = model_names config._save_models_to_file() return {"models": config.model_names[model_provider]["models"]} + @chat.get("/tools") async def get_tools(agent_id: str, current_user: User = Depends(get_required_user)): """获取所有可用工具(需要登录)""" @@ -193,12 +203,9 @@ async def get_tools(agent_id: str, current_user: User = Depends(get_required_use tools_info = gen_tool_info(tools) return {"tools": {tool["id"]: tool for tool in tools_info}} + @chat.post("/agent/{agent_id}/config") -async def save_agent_config( - agent_id: str, - config: dict = Body(...), - current_user: User = Depends(get_required_user) -): +async def save_agent_config(agent_id: str, config: dict = Body(...), current_user: User = Depends(get_required_user)): """保存智能体配置到YAML文件(需要登录)""" try: # 获取Agent实例和配置类 @@ -217,12 +224,9 @@ async def save_agent_config( logger.error(f"保存智能体配置出错: {e}, {traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"保存智能体配置出错: {str(e)}") + @chat.get("/agent/{agent_id}/history") -async def get_agent_history( - agent_id: str, - thread_id: str, - current_user: User = Depends(get_required_user) -): +async def get_agent_history(agent_id: str, thread_id: str, current_user: User = Depends(get_required_user)): """获取智能体历史消息(需要登录)""" try: # 获取Agent实例和配置类 @@ -237,11 +241,9 @@ async def get_agent_history( logger.error(f"获取智能体历史消息出错: {e}, {traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"获取智能体历史消息出错: {str(e)}") + @chat.get("/agent/{agent_id}/config") -async def get_agent_config( - agent_id: str, - current_user: User = Depends(get_required_user) -): +async def get_agent_config(agent_id: str, current_user: User = Depends(get_required_user)): """从YAML文件加载智能体配置(需要登录)""" try: # 检查智能体是否存在 @@ -256,8 +258,10 @@ async def get_agent_config( logger.error(f"加载智能体配置出错: {e}, {traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"加载智能体配置出错: {str(e)}") + # ==================== 线程管理 API ==================== + class ThreadCreate(BaseModel): title: str | None = None agent_id: str @@ -279,11 +283,10 @@ class ThreadResponse(BaseModel): # > === 会话管理分组 === # ============================================================================= + @chat.post("/thread", response_model=ThreadResponse) async def create_thread( - thread: ThreadCreate, - db: Session = Depends(get_db), - current_user: User = Depends(get_required_user) + thread: ThreadCreate, db: Session = Depends(get_db), current_user: User = Depends(get_required_user) ): """创建新对话线程""" thread_id = str(uuid.uuid4()) @@ -313,11 +316,7 @@ async def create_thread( @chat.get("/threads", response_model=list[ThreadResponse]) -async def list_threads( - agent_id: str, - db: Session = Depends(get_db), - current_user: User = Depends(get_required_user) -): +async def list_threads(agent_id: str, db: Session = Depends(get_db), current_user: User = Depends(get_required_user)): """获取用户的所有对话线程""" assert agent_id, "agent_id 不能为空" query = db.query(Thread).filter( @@ -345,16 +344,9 @@ async def list_threads( @chat.delete("/thread/{thread_id}") -async def delete_thread( - thread_id: str, - db: Session = Depends(get_db), - current_user: User = Depends(get_required_user) -): +async def delete_thread(thread_id: str, db: Session = Depends(get_db), current_user: User = Depends(get_required_user)): """删除对话线程""" - thread = db.query(Thread).filter( - Thread.id == thread_id, - Thread.user_id == str(current_user.id) - ).first() + thread = db.query(Thread).filter(Thread.id == thread_id, Thread.user_id == str(current_user.id)).first() if not thread: raise HTTPException(status_code=404, detail="对话线程不存在") @@ -376,14 +368,14 @@ async def update_thread( thread_id: str, thread_update: ThreadUpdate, db: Session = Depends(get_db), - current_user: User = Depends(get_required_user) + current_user: User = Depends(get_required_user), ): """更新对话线程信息""" - thread = db.query(Thread).filter( - Thread.id == thread_id, - Thread.user_id == str(current_user.id), - Thread.status == 1 - ).first() + thread = ( + db.query(Thread) + .filter(Thread.id == thread_id, Thread.user_id == str(current_user.id), Thread.status == 1) + .first() + ) if not thread: raise HTTPException(status_code=404, detail="对话线程不存在") diff --git a/server/routers/graph_router.py b/server/routers/graph_router.py index 0b348e38..04df5bb0 100644 --- a/server/routers/graph_router.py +++ b/server/routers/graph_router.py @@ -1,9 +1,10 @@ import traceback -from fastapi import APIRouter, Query, HTTPException, Depends, Body -from server.utils.auth_middleware import get_admin_user -from server.models.user_model import User -from src import knowledge_base, graph_base +from fastapi import APIRouter, Body, Depends, HTTPException, Query + +from server.models.user_model import User +from server.utils.auth_middleware import get_admin_user +from src import graph_base, knowledge_base from src.utils.logging_config import logger graph = APIRouter(prefix="/graph", tags=["graph"]) @@ -13,13 +14,14 @@ graph = APIRouter(prefix="/graph", tags=["graph"]) # === 子图查询分组 === # ============================================================================= + @graph.get("/lightrag/subgraph") async def get_lightrag_subgraph( db_id: str = Query(..., description="数据库ID"), node_label: str = Query(..., description="节点标签或实体名称"), max_depth: int = Query(2, description="最大深度", ge=1, le=5), max_nodes: int = Query(100, description="最大节点数", ge=1, le=1000), - current_user: User = Depends(get_admin_user) + current_user: User = Depends(get_admin_user), ): """ 使用 LightRAG 原生方法获取知识图谱子图 @@ -34,13 +36,14 @@ async def get_lightrag_subgraph( 包含节点和边的知识图谱数据 """ try: - logger.info(f"获取子图数据 - db_id: {db_id}, node_label: {node_label}, max_depth: {max_depth}, max_nodes: {max_nodes}") + logger.info( + f"获取子图数据 - db_id: {db_id}, node_label: {node_label}, max_depth: {max_depth}, max_nodes: {max_nodes}" + ) # 检查是否是 LightRAG 数据库 if not knowledge_base.is_lightrag_database(db_id): raise HTTPException( - status_code=400, - detail=f"数据库 {db_id} 不是 LightRAG 类型,图谱功能仅支持 LightRAG 知识库" + status_code=400, detail=f"数据库 {db_id} 不是 LightRAG 类型,图谱功能仅支持 LightRAG 知识库" ) # 获取 LightRAG 实例 @@ -50,30 +53,32 @@ async def get_lightrag_subgraph( # 使用 LightRAG 的原生 get_knowledge_graph 方法 knowledge_graph = await rag_instance.get_knowledge_graph( - node_label=node_label, - max_depth=max_depth, - max_nodes=max_nodes + node_label=node_label, max_depth=max_depth, max_nodes=max_nodes ) # 将 LightRAG 的 KnowledgeGraph 格式转换为前端需要的格式 nodes = [] for node in knowledge_graph.nodes: - nodes.append({ - "id": node.id, - "labels": node.labels, - "entity_type": node.properties.get("entity_type", "unknown"), - "properties": node.properties - }) + nodes.append( + { + "id": node.id, + "labels": node.labels, + "entity_type": node.properties.get("entity_type", "unknown"), + "properties": node.properties, + } + ) edges = [] for edge in knowledge_graph.edges: - edges.append({ - "id": edge.id, - "source": edge.source, - "target": edge.target, - "type": edge.type, - "properties": edge.properties - }) + edges.append( + { + "id": edge.id, + "source": edge.source, + "target": edge.target, + "type": edge.type, + "properties": edge.properties, + } + ) result = { "success": True, @@ -82,8 +87,8 @@ async def get_lightrag_subgraph( "edges": edges, "is_truncated": knowledge_graph.is_truncated, "total_nodes": len(nodes), - "total_edges": len(edges) - } + "total_edges": len(edges), + }, } logger.info(f"成功获取子图 - 节点数: {len(nodes)}, 边数: {len(edges)}") @@ -99,9 +104,7 @@ async def get_lightrag_subgraph( @graph.get("/lightrag/databases") -async def get_lightrag_databases( - current_user: User = Depends(get_admin_user) -): +async def get_lightrag_databases(current_user: User = Depends(get_admin_user)): """ 获取所有可用的 LightRAG 数据库 @@ -110,25 +113,21 @@ async def get_lightrag_databases( """ try: lightrag_databases = knowledge_base.get_lightrag_databases() - return { - "success": True, - "data": { - "databases": lightrag_databases - } - } + return {"success": True, "data": {"databases": lightrag_databases}} except Exception as e: logger.error(f"获取 LightRAG 数据库列表失败: {e}") raise HTTPException(status_code=500, detail=f"获取 LightRAG 数据库列表失败: {str(e)}") + # ============================================================================= # === 节点管理分组 === # ============================================================================= + @graph.get("/lightrag/labels") async def get_lightrag_labels( - db_id: str = Query(..., description="数据库ID"), - current_user: User = Depends(get_admin_user) + db_id: str = Query(..., description="数据库ID"), current_user: User = Depends(get_admin_user) ): """ 获取知识图谱中的所有标签 @@ -145,8 +144,7 @@ async def get_lightrag_labels( # 检查是否是 LightRAG 数据库 if not knowledge_base.is_lightrag_database(db_id): raise HTTPException( - status_code=400, - detail=f"数据库 {db_id} 不是 LightRAG 类型,图谱功能仅支持 LightRAG 知识库" + status_code=400, detail=f"数据库 {db_id} 不是 LightRAG 类型,图谱功能仅支持 LightRAG 知识库" ) # 获取 LightRAG 实例 @@ -157,12 +155,7 @@ async def get_lightrag_labels( # 使用 LightRAG 的原生方法获取所有标签 labels = await rag_instance.get_graph_labels() - return { - "success": True, - "data": { - "labels": labels - } - } + return {"success": True, "data": {"labels": labels}} except HTTPException: # 重新抛出 HTTP 异常 @@ -177,7 +170,7 @@ async def get_lightrag_labels( async def get_neo4j_nodes( kgdb_name: str = Query(..., description="知识图谱数据库名称"), num: int = Query(100, description="节点数量", ge=1, le=1000), - current_user: User = Depends(get_admin_user) + current_user: User = Depends(get_admin_user), ): """ 获取图谱节点样本数据 @@ -190,20 +183,16 @@ async def get_neo4j_nodes( result = graph_base.get_sample_nodes(kgdb_name, num) - return { - "success": True, - "result": result, - "message": "success" - } + return {"success": True, "result": result, "message": "success"} except Exception as e: logger.error(f"获取图节点数据失败: {e}\n{traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"获取图节点数据失败: {str(e)}") + @graph.get("/neo4j/node") async def get_neo4j_node( - entity_name: str = Query(..., description="实体名称"), - current_user: User = Depends(get_admin_user) + entity_name: str = Query(..., description="实体名称"), current_user: User = Depends(get_admin_user) ): """ 根据实体名称查询图节点 @@ -214,16 +203,13 @@ async def get_neo4j_node( result = graph_base.query_node(entity_name=entity_name) - return { - "success": True, - "result": result, - "message": "success" - } + return {"success": True, "result": result, "message": "success"} except Exception as e: logger.error(f"查询图节点失败: {e}\n{traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"查询图节点失败: {str(e)}") + # ============================================================================= # === 边管理分组 === # ============================================================================= @@ -234,10 +220,10 @@ async def get_neo4j_node( # === 图谱分析分组 === # ============================================================================= + @graph.get("/lightrag/stats") async def get_lightrag_stats( - db_id: str = Query(..., description="数据库ID"), - current_user: User = Depends(get_admin_user) + db_id: str = Query(..., description="数据库ID"), current_user: User = Depends(get_admin_user) ): """ 获取知识图谱统计信息 @@ -248,8 +234,7 @@ async def get_lightrag_stats( # 检查是否是 LightRAG 数据库 if not knowledge_base.is_lightrag_database(db_id): raise HTTPException( - status_code=400, - detail=f"数据库 {db_id} 不是 LightRAG 类型,图谱功能仅支持 LightRAG 知识库" + status_code=400, detail=f"数据库 {db_id} 不是 LightRAG 类型,图谱功能仅支持 LightRAG 知识库" ) # 获取 LightRAG 实例 @@ -261,7 +246,7 @@ async def get_lightrag_stats( knowledge_graph = await rag_instance.get_knowledge_graph( node_label="*", max_depth=1, - max_nodes=10000 # 设置较大值以获取完整统计 + max_nodes=10000, # 设置较大值以获取完整统计 ) # 统计实体类型分布 @@ -271,8 +256,7 @@ async def get_lightrag_stats( entity_types[entity_type] = entity_types.get(entity_type, 0) + 1 entity_types_list = [ - {"type": k, "count": v} - for k, v in sorted(entity_types.items(), key=lambda x: x[1], reverse=True) + {"type": k, "count": v} for k, v in sorted(entity_types.items(), key=lambda x: x[1], reverse=True) ] return { @@ -281,8 +265,8 @@ async def get_lightrag_stats( "total_nodes": len(knowledge_graph.nodes), "total_edges": len(knowledge_graph.edges), "entity_types": entity_types_list, - "is_truncated": knowledge_graph.is_truncated - } + "is_truncated": knowledge_graph.is_truncated, + }, } except HTTPException: @@ -293,6 +277,7 @@ async def get_lightrag_stats( logger.error(f"Traceback: {traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"获取图谱统计信息失败: {str(e)}") + @graph.get("/neo4j/info") async def get_neo4j_info(current_user: User = Depends(get_admin_user)): """获取Neo4j图数据库信息""" @@ -300,26 +285,21 @@ async def get_neo4j_info(current_user: User = Depends(get_admin_user)): graph_info = graph_base.get_graph_info() if graph_info is None: raise HTTPException(status_code=400, detail="图数据库获取出错") - return { - "success": True, - "data": graph_info - } + return {"success": True, "data": graph_info} except Exception as e: logger.error(f"获取图数据库信息失败: {e}") raise HTTPException(status_code=500, detail=f"获取图数据库信息失败: {str(e)}") + @graph.post("/neo4j/index-entities") -async def index_neo4j_entities( - data: dict = Body(default={}), - current_user: User = Depends(get_admin_user) -): +async def index_neo4j_entities(data: dict = Body(default={}), current_user: User = Depends(get_admin_user)): """为Neo4j图谱节点添加嵌入向量索引""" try: if not graph_base.is_running(): raise HTTPException(status_code=400, detail="图数据库未启动") # 获取参数或使用默认值 - kgdb_name = data.get('kgdb_name', 'neo4j') + kgdb_name = data.get("kgdb_name", "neo4j") # 调用GraphDatabase的add_embedding_to_nodes方法 count = graph_base.add_embedding_to_nodes(kgdb_name=kgdb_name) @@ -328,37 +308,24 @@ async def index_neo4j_entities( "success": True, "status": "success", "message": f"已成功为{count}个节点添加嵌入向量", - "indexed_count": count + "indexed_count": count, } except Exception as e: logger.error(f"索引节点失败: {e}") raise HTTPException(status_code=500, detail=f"索引节点失败: {str(e)}") + @graph.post("/neo4j/add-entities") async def add_neo4j_entities( - file_path: str = Body(...), - kgdb_name: str | None = Body(None), - current_user: User = Depends(get_admin_user) + file_path: str = Body(...), kgdb_name: str | None = Body(None), current_user: User = Depends(get_admin_user) ): """通过JSONL文件添加图谱实体到Neo4j""" try: - if not file_path.endswith('.jsonl'): - return { - "success": False, - "message": "文件格式错误,请上传jsonl文件", - "status": "failed" - } + if not file_path.endswith(".jsonl"): + return {"success": False, "message": "文件格式错误,请上传jsonl文件", "status": "failed"} await graph_base.jsonl_file_add_entity(file_path, kgdb_name) - return { - "success": True, - "message": "实体添加成功", - "status": "success" - } + return {"success": True, "message": "实体添加成功", "status": "success"} except Exception as e: logger.error(f"添加实体失败: {e}, {traceback.format_exc()}") - return { - "success": False, - "message": f"添加实体失败: {e}", - "status": "failed" - } + return {"success": False, "message": f"添加实体失败: {e}", "status": "failed"} diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index 68baa99d..168a9158 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -1,14 +1,15 @@ -import os import asyncio +import os import traceback -from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body, Form, Query + +from fastapi import APIRouter, Body, Depends, File, Form, HTTPException, Query, UploadFile from fastapi.responses import FileResponse -from src import executor, config, knowledge_base -from src.utils import logger, hashstr -from src.knowledge.indexing import process_file_to_markdown -from server.utils.auth_middleware import get_admin_user from server.models.user_model import User +from server.utils.auth_middleware import get_admin_user +from src import config, executor, knowledge_base +from src.knowledge.indexing import process_file_to_markdown +from src.utils import hashstr, logger knowledge = APIRouter(prefix="/knowledge", tags=["knowledge"]) @@ -16,6 +17,7 @@ knowledge = APIRouter(prefix="/knowledge", tags=["knowledge"]) # === 数据库管理分组 === # ============================================================================= + @knowledge.get("/databases") async def get_databases(current_user: User = Depends(get_admin_user)): """获取所有知识库""" @@ -26,6 +28,7 @@ async def get_databases(current_user: User = Depends(get_admin_user)): logger.error(f"获取数据库列表失败 {e}, {traceback.format_exc()}") return {"message": f"获取数据库列表失败 {e}", "databases": []} + @knowledge.post("/databases") async def create_database( database_name: str = Body(...), @@ -34,23 +37,21 @@ async def create_database( kb_type: str = Body("lightrag"), additional_params: dict = Body({}), llm_info: dict = Body(None), - current_user: User = Depends(get_admin_user) + current_user: User = Depends(get_admin_user), ): """创建知识库""" - logger.debug(f"Create database {database_name} with kb_type {kb_type}, additional_params {additional_params}, llm_info {llm_info}") + logger.debug( + f"Create database {database_name} with kb_type {kb_type}, additional_params {additional_params}, llm_info {llm_info}" + ) try: embed_info = config.embed_model_names[embed_model_name] database_info = await knowledge_base.create_database( - database_name, - description, - kb_type=kb_type, - embed_info=embed_info, - llm_info=llm_info, - **additional_params + database_name, description, kb_type=kb_type, embed_info=embed_info, llm_info=llm_info, **additional_params ) # 需要重新加载所有智能体,因为工具刷新了 from src.agents import agent_manager + await agent_manager.reload_all() return database_info @@ -58,6 +59,7 @@ async def create_database( logger.error(f"创建数据库失败 {e}, {traceback.format_exc()}") return {"message": f"创建数据库失败 {e}", "status": "failed"} + @knowledge.get("/databases/{db_id}") async def get_database_info(db_id: str, current_user: User = Depends(get_admin_user)): """获取知识库详细信息""" @@ -66,12 +68,10 @@ async def get_database_info(db_id: str, current_user: User = Depends(get_admin_u raise HTTPException(status_code=404, detail="Database not found") return database + @knowledge.put("/databases/{db_id}") async def update_database_info( - db_id: str, - name: str = Body(...), - description: str = Body(...), - current_user: User = Depends(get_admin_user) + db_id: str, name: str = Body(...), description: str = Body(...), current_user: User = Depends(get_admin_user) ): """更新知识库信息""" logger.debug(f"Update database {db_id} info: {name}, {description}") @@ -82,6 +82,7 @@ async def update_database_info( logger.error(f"更新数据库失败 {e}, {traceback.format_exc()}") raise HTTPException(status_code=400, detail=f"更新数据库失败: {e}") + @knowledge.delete("/databases/{db_id}") async def delete_database(db_id: str, current_user: User = Depends(get_admin_user)): """删除知识库""" @@ -91,6 +92,7 @@ async def delete_database(db_id: str, current_user: User = Depends(get_admin_use # 需要重新加载所有智能体,因为工具刷新了 from src.agents import agent_manager + await agent_manager.reload_all() return {"message": "删除成功"} @@ -98,19 +100,18 @@ async def delete_database(db_id: str, current_user: User = Depends(get_admin_use logger.error(f"删除数据库失败 {e}, {traceback.format_exc()}") raise HTTPException(status_code=400, detail=f"删除数据库失败: {e}") + @knowledge.get("/databases/{db_id}/export") async def export_database( db_id: str, format: str = Query("csv", enum=["csv", "xlsx", "md", "txt"]), include_vectors: bool = Query(False, description="是否在导出中包含向量数据"), - current_user: User = Depends(get_admin_user) + current_user: User = Depends(get_admin_user), ): """导出知识库数据""" logger.debug(f"Exporting database {db_id} with format {format}") try: - file_path = await knowledge_base.export_data( - db_id, format=format, include_vectors=include_vectors - ) + file_path = await knowledge_base.export_data(db_id, format=format, include_vectors=include_vectors) if not os.path.exists(file_path): raise HTTPException(status_code=404, detail="Exported file not found.") @@ -119,15 +120,11 @@ async def export_database( "csv": "text/csv", "xlsx": "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", "md": "text/markdown", - "txt": "text/plain" + "txt": "text/plain", } media_type = media_types.get(format, "application/octet-stream") - return FileResponse( - path=file_path, - filename=os.path.basename(file_path), - media_type=media_type - ) + return FileResponse(path=file_path, filename=os.path.basename(file_path), media_type=media_type) except NotImplementedError as e: logger.warning(f"A disabled feature was accessed: {e}") raise HTTPException(status_code=501, detail=str(e)) @@ -135,38 +132,34 @@ async def export_database( logger.error(f"导出数据库失败 {e}, {traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"导出数据库失败: {e}") + # ============================================================================= # === 文档管理分组 === # ============================================================================= + @knowledge.post("/databases/{db_id}/documents") async def add_documents( - db_id: str, - items: list[str] = Body(...), - params: dict = Body(...), - current_user: User = Depends(get_admin_user) + db_id: str, items: list[str] = Body(...), params: dict = Body(...), current_user: User = Depends(get_admin_user) ): """添加文档到知识库""" logger.debug(f"Add documents for db_id {db_id}: {items} {params=}") - content_type = params.get('content_type', 'file') + content_type = params.get("content_type", "file") try: processed_items = await knowledge_base.add_content(db_id, items, params=params) - item_type = "URLs" if content_type == 'url' else "files" - processed_failed_count = len([_p for _p in processed_items if _p['status'] == 'failed']) + item_type = "URLs" if content_type == "url" else "files" + processed_failed_count = len([_p for _p in processed_items if _p["status"] == "failed"]) processed_info = f"Processed {len(processed_items)} {item_type}, {processed_failed_count} {item_type} failed" return {"message": processed_info, "items": processed_items, "status": "success"} except Exception as e: logger.error(f"Failed to process {content_type}s: {e}, {traceback.format_exc()}") return {"message": f"Failed to process {content_type}s: {e}", "status": "failed"} + @knowledge.get("/databases/{db_id}/documents/{doc_id}") -async def get_document_info( - db_id: str, - doc_id: str, - current_user: User = Depends(get_admin_user) -): +async def get_document_info(db_id: str, doc_id: str, current_user: User = Depends(get_admin_user)): """获取文档详细信息""" logger.debug(f"GET document {doc_id} info in {db_id}") @@ -177,12 +170,9 @@ async def get_document_info( logger.error(f"Failed to get file info, {e}, {db_id=}, {doc_id=}, {traceback.format_exc()}") return {"message": "Failed to get file info", "status": "failed"} + @knowledge.delete("/databases/{db_id}/documents/{doc_id}") -async def delete_document( - db_id: str, - doc_id: str, - current_user: User = Depends(get_admin_user) -): +async def delete_document(db_id: str, doc_id: str, current_user: User = Depends(get_admin_user)): """删除文档""" logger.debug(f"DELETE document {doc_id} info in {db_id}") try: @@ -192,16 +182,15 @@ async def delete_document( logger.error(f"删除文档失败 {e}, {traceback.format_exc()}") raise HTTPException(status_code=400, detail=f"删除文档失败: {e}") + # ============================================================================= # === 查询分组 === # ============================================================================= + @knowledge.post("/databases/{db_id}/query") async def query_knowledge_base( - db_id: str, - query: str = Body(...), - meta: dict = Body(...), - current_user: User = Depends(get_admin_user) + db_id: str, query: str = Body(...), meta: dict = Body(...), current_user: User = Depends(get_admin_user) ): """查询知识库""" logger.debug(f"Query knowledge base {db_id}: {query}") @@ -212,12 +201,10 @@ async def query_knowledge_base( logger.error(f"知识库查询失败 {e}, {traceback.format_exc()}") return {"message": f"知识库查询失败: {e}", "status": "failed"} + @knowledge.post("/databases/{db_id}/query-test") async def query_test( - db_id: str, - query: str = Body(...), - meta: dict = Body(...), - current_user: User = Depends(get_admin_user) + db_id: str, query: str = Body(...), meta: dict = Body(...), current_user: User = Depends(get_admin_user) ): """测试查询知识库""" logger.debug(f"Query test in {db_id}: {query}") @@ -228,11 +215,9 @@ async def query_test( logger.error(f"测试查询失败 {e}, {traceback.format_exc()}") return {"message": f"测试查询失败: {e}", "status": "failed"} + @knowledge.get("/databases/{db_id}/query-params") -async def get_knowledge_base_query_params( - db_id: str, - current_user: User = Depends(get_admin_user) -): +async def get_knowledge_base_query_params(db_id: str, current_user: User = Depends(get_admin_user)): """获取知识库类型特定的查询参数""" try: # 获取数据库信息 @@ -258,21 +243,21 @@ async def get_knowledge_base_query_params( {"value": "hybrid", "label": "Hybrid", "description": "本地和全局混合"}, {"value": "naive", "label": "Naive", "description": "基本搜索"}, {"value": "mix", "label": "Mix", "description": "知识图谱和向量检索混合"}, - ] + ], }, { "key": "only_need_context", "label": "只使用上下文", "type": "boolean", "default": True, - "description": "只返回上下文,不生成回答" + "description": "只返回上下文,不生成回答", }, { "key": "only_need_prompt", "label": "只使用提示", "type": "boolean", "default": False, - "description": "只返回提示,不进行检索" + "description": "只返回提示,不进行检索", }, { "key": "top_k", @@ -281,9 +266,9 @@ async def get_knowledge_base_query_params( "default": 10, "min": 1, "max": 100, - "description": "返回的最大结果数量" - } - ] + "description": "返回的最大结果数量", + }, + ], } elif kb_type == "chroma": params = { @@ -296,7 +281,7 @@ async def get_knowledge_base_query_params( "default": 10, "min": 1, "max": 100, - "description": "返回的最大结果数量" + "description": "返回的最大结果数量", }, { "key": "similarity_threshold", @@ -306,16 +291,16 @@ async def get_knowledge_base_query_params( "min": 0.0, "max": 1.0, "step": 0.1, - "description": "过滤相似度低于此值的结果" + "description": "过滤相似度低于此值的结果", }, { "key": "include_distances", "label": "显示相似度", "type": "boolean", "default": True, - "description": "在结果中显示相似度分数" - } - ] + "description": "在结果中显示相似度分数", + }, + ], } elif kb_type == "milvus": params = { @@ -328,7 +313,7 @@ async def get_knowledge_base_query_params( "default": 10, "min": 1, "max": 100, - "description": "返回的最大结果数量" + "description": "返回的最大结果数量", }, { "key": "similarity_threshold", @@ -338,14 +323,14 @@ async def get_knowledge_base_query_params( "min": 0.0, "max": 1.0, "step": 0.1, - "description": "过滤相似度低于此值的结果" + "description": "过滤相似度低于此值的结果", }, { "key": "include_distances", "label": "显示相似度", "type": "boolean", "default": True, - "description": "在结果中显示相似度分数" + "description": "在结果中显示相似度分数", }, { "key": "metric_type", @@ -355,11 +340,11 @@ async def get_knowledge_base_query_params( "options": [ {"value": "COSINE", "label": "余弦相似度", "description": "适合文本语义相似度"}, {"value": "L2", "label": "欧几里得距离", "description": "适合数值型数据"}, - {"value": "IP", "label": "内积", "description": "适合标准化向量"} + {"value": "IP", "label": "内积", "description": "适合标准化向量"}, ], - "description": "向量相似度计算方法" - } - ] + "description": "向量相似度计算方法", + }, + ], } else: # 未知类型,返回基本参数 @@ -373,9 +358,9 @@ async def get_knowledge_base_query_params( "default": 10, "min": 1, "max": 100, - "description": "返回的最大结果数量" + "description": "返回的最大结果数量", } - ] + ], } return {"params": params, "message": "success"} @@ -384,15 +369,15 @@ async def get_knowledge_base_query_params( logger.error(f"获取知识库查询参数失败 {e}, {traceback.format_exc()}") return {"message": f"获取知识库查询参数失败 {e}", "params": {}} + # ============================================================================= # === 文件管理分组 === # ============================================================================= + @knowledge.post("/files/upload") async def upload_file( - file: UploadFile = File(...), - db_id: str | None = Query(None), - current_user: User = Depends(get_admin_user) + file: UploadFile = File(...), db_id: str | None = Query(None), current_user: User = Depends(get_admin_user) ): """上传文件""" if not file.filename: @@ -414,11 +399,9 @@ async def upload_file( return {"message": "File successfully uploaded", "file_path": file_path, "db_id": db_id} + @knowledge.post("/files/markdown") -async def mark_it_down( - file: UploadFile = File(...), - current_user: User = Depends(get_admin_user) -): +async def mark_it_down(file: UploadFile = File(...), current_user: User = Depends(get_admin_user)): """调用 src.knowledge.indexing 下面的 process_file_to_markdown 解析为 markdown,参数是文件,需要管理员权限""" try: content = await file.read() @@ -428,10 +411,12 @@ async def mark_it_down( logger.error(f"文件解析失败 {e}, {traceback.format_exc()}") return {"message": f"文件解析失败 {e}", "markdown_content": ""} + # ============================================================================= # === 知识库类型分组 === # ============================================================================= + @knowledge.get("/types") async def get_knowledge_base_types(current_user: User = Depends(get_admin_user)): """获取支持的知识库类型""" @@ -442,6 +427,7 @@ async def get_knowledge_base_types(current_user: User = Depends(get_admin_user)) logger.error(f"获取知识库类型失败 {e}, {traceback.format_exc()}") return {"message": f"获取知识库类型失败 {e}", "kb_types": {}} + @knowledge.get("/stats") async def get_knowledge_base_statistics(current_user: User = Depends(get_admin_user)): """获取知识库统计信息""" diff --git a/server/routers/system_router.py b/server/routers/system_router.py index a4925659..f7945346 100644 --- a/server/routers/system_router.py +++ b/server/routers/system_router.py @@ -1,64 +1,63 @@ import os -import yaml -import requests -from pathlib import Path -from fastapi import Request, Body, Depends, HTTPException -from fastapi import APIRouter from collections import deque +from pathlib import Path + +import requests +import yaml +from fastapi import APIRouter, Body, Depends, HTTPException, Request -from src import config, knowledge_base, graph_base -from server.utils.auth_middleware import get_admin_user, get_superadmin_user from server.models.user_model import User +from server.utils.auth_middleware import get_admin_user, get_superadmin_user +from src import config, graph_base, knowledge_base from src.utils.logging_config import logger - system = APIRouter(prefix="/system", tags=["system"]) # ============================================================================= # === 健康检查分组 === # ============================================================================= + @system.get("/health") async def health_check(): """系统健康检查接口(公开接口)""" return {"status": "ok", "message": "服务正常运行"} + # ============================================================================= # === 配置管理分组 === # ============================================================================= + @system.get("/config") def get_config(current_user: User = Depends(get_admin_user)): """获取系统配置""" return config.dump_config() + @system.post("/config") -async def update_config_single( - key = Body(...), - value = Body(...), - current_user: User = Depends(get_admin_user) -) -> dict: +async def update_config_single(key=Body(...), value=Body(...), current_user: User = Depends(get_admin_user)) -> dict: """更新单个配置项""" config[key] = value config.save() return config.dump_config() + @system.post("/config/update") -async def update_config_batch( - items: dict = Body(...), - current_user: User = Depends(get_admin_user) -) -> dict: +async def update_config_batch(items: dict = Body(...), current_user: User = Depends(get_admin_user)) -> dict: """批量更新配置项""" config.update(items) config.save() return config.dump_config() + @system.post("/restart") async def restart_system(current_user: User = Depends(get_superadmin_user)): """重启系统(仅超级管理员)""" graph_base.start() return {"message": "系统已重启"} + @system.get("/logs") def get_system_logs(current_user: User = Depends(get_admin_user)): """获取系统日志""" @@ -68,16 +67,18 @@ def get_system_logs(current_user: User = Depends(get_admin_user)): with open(LOG_FILE) as f: last_lines = deque(f, maxlen=1000) - log = ''.join(last_lines) + log = "".join(last_lines) return {"log": log, "message": "success", "log_file": LOG_FILE} except Exception as e: logger.error(f"获取系统日志失败: {e}") raise HTTPException(status_code=500, detail=f"获取系统日志失败: {str(e)}") + # ============================================================================= # === 信息管理分组 === # ============================================================================= + def load_info_config(): """加载信息配置文件""" try: @@ -91,7 +92,7 @@ def load_info_config(): config_path = Path("src/static/info.template.yaml") # 读取配置文件 - with open(config_path, encoding='utf-8') as file: + with open(config_path, encoding="utf-8") as file: config = yaml.safe_load(file) return config @@ -100,61 +101,49 @@ def load_info_config(): logger.error(f"Failed to load info config: {e}") return get_default_info_config() + def get_default_info_config(): """获取默认信息配置""" return { - "organization": { - "name": "江南语析", - "logo": "/favicon.svg", - "avatar": "/avatar.jpg" - }, + "organization": {"name": "江南语析", "logo": "/favicon.svg", "avatar": "/avatar.jpg"}, "branding": { "name": "Yuxi-Know", "title": "Yuxi-Know", "subtitle": "大模型驱动的知识库管理工具", - "description": "结合知识库与知识图谱,提供更准确、更全面的回答" + "description": "结合知识库与知识图谱,提供更准确、更全面的回答", }, - "features": [ - "📚 灵活知识库", - "🕸️ 知识图谱集成", - "🤖 多模型支持" - ], - "footer": { - "copyright": "© 江南语析 2025 [WIP] v0.2.0" - } + "features": ["📚 灵活知识库", "🕸️ 知识图谱集成", "🤖 多模型支持"], + "footer": {"copyright": "© 江南语析 2025 [WIP] v0.2.0"}, } + @system.get("/info") async def get_info_config(): """获取系统信息配置(公开接口,无需认证)""" try: config = load_info_config() - return { - "success": True, - "data": config - } + return {"success": True, "data": config} except Exception as e: logger.error(f"获取信息配置失败: {e}") raise HTTPException(status_code=500, detail="获取信息配置失败") + @system.post("/info/reload") async def reload_info_config(current_user: User = Depends(get_admin_user)): """重新加载信息配置""" try: config = load_info_config() - return { - "success": True, - "message": "配置重新加载成功", - "data": config - } + return {"success": True, "message": "配置重新加载成功", "data": config} except Exception as e: logger.error(f"重新加载信息配置失败: {e}") raise HTTPException(status_code=500, detail="重新加载信息配置失败") + # ============================================================================= # === OCR服务分组 === # ============================================================================= + @system.get("/ocr/stats") async def get_ocr_stats(current_user: User = Depends(get_admin_user)): """ @@ -163,20 +152,13 @@ async def get_ocr_stats(current_user: User = Depends(get_admin_user)): """ try: from src.plugins._ocr import get_ocr_stats + stats = get_ocr_stats() - return { - "status": "success", - "stats": stats, - "message": "OCR统计信息获取成功" - } + return {"status": "success", "stats": stats, "message": "OCR统计信息获取成功"} except Exception as e: logger.error(f"获取OCR统计信息失败: {str(e)}") - return { - "status": "error", - "stats": {}, - "message": f"获取OCR统计信息失败: {str(e)}" - } + return {"status": "error", "stats": {}, "message": f"获取OCR统计信息失败: {str(e)}"} @system.get("/ocr/health") @@ -188,12 +170,14 @@ async def check_ocr_services_health(current_user: User = Depends(get_admin_user) health_status = { "rapid_ocr": {"status": "unknown", "message": ""}, "mineru_ocr": {"status": "unknown", "message": ""}, - "paddlex_ocr": {"status": "unknown", "message": ""} + "paddlex_ocr": {"status": "unknown", "message": ""}, } # 检查 RapidOCR (ONNX) 模型 try: - model_dir_root = os.getenv("MODEL_DIR") if not os.getenv("RUNNING_IN_DOCKER") else os.getenv("MODEL_DIR_IN_DOCKER") + model_dir_root = ( + os.getenv("MODEL_DIR") if not os.getenv("RUNNING_IN_DOCKER") else os.getenv("MODEL_DIR_IN_DOCKER") + ) model_dir = os.path.join(model_dir_root, "SWHL/RapidOCR") det_model_path = os.path.join(model_dir, "PP-OCRv4/ch_PP-OCRv4_det_infer.onnx") rec_model_path = os.path.join(model_dir, "PP-OCRv4/ch_PP-OCRv4_rec_infer.onnx") @@ -201,6 +185,7 @@ async def check_ocr_services_health(current_user: User = Depends(get_admin_user) if os.path.exists(model_dir) and os.path.exists(det_model_path) and os.path.exists(rec_model_path): # 尝试初始化RapidOCR from rapidocr_onnxruntime import RapidOCR + test_ocr = RapidOCR(det_box_thresh=0.3, det_model_path=det_model_path, rec_model_path=rec_model_path) # noqa: F841 health_status["rapid_ocr"]["status"] = "healthy" health_status["rapid_ocr"]["message"] = "RapidOCR模型已加载" @@ -258,8 +243,4 @@ async def check_ocr_services_health(current_user: User = Depends(get_admin_user) # 计算整体健康状态 overall_status = "healthy" if any(svc["status"] == "healthy" for svc in health_status.values()) else "unhealthy" - return { - "overall_status": overall_status, - "services": health_status, - "message": "OCR服务健康检查完成" - } + return {"overall_status": overall_status, "services": health_status, "message": "OCR服务健康检查完成"} diff --git a/server/utils/auth_middleware.py b/server/utils/auth_middleware.py index 346d987f..57920f49 100644 --- a/server/utils/auth_middleware.py +++ b/server/utils/auth_middleware.py @@ -1,8 +1,9 @@ +import re + from fastapi import Depends, HTTPException, status from fastapi.security import OAuth2PasswordBearer -from sqlalchemy.orm import Session from jose import JWTError, jwt -import re +from sqlalchemy.orm import Session from server.db_manager import db_manager from server.models.user_model import User @@ -13,14 +14,15 @@ oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/token", auto_error=Fals # 公开路径列表,无需登录即可访问 PUBLIC_PATHS = [ - r"^/api/auth/token$", # 登录 + r"^/api/auth/token$", # 登录 r"^/api/auth/check-first-run$", # 检查是否首次运行 - r"^/api/auth/initialize$", # 初始化系统 - r"^/api$", # Health Check - r"^/api/system/health$", # Health Check - r"^/api/system/info$", # 获取系统信息配置 + r"^/api/auth/initialize$", # 初始化系统 + r"^/api$", # Health Check + r"^/api/system/health$", # Health Check + r"^/api/system/info$", # 获取系统信息配置 ] + # 获取数据库会话 def get_db(): db = db_manager.get_session() @@ -29,6 +31,7 @@ def get_db(): finally: db.close() + # 获取当前用户 async def get_current_user(token: str | None = Depends(oauth2_scheme), db: Session = Depends(get_db)): credentials_exception = HTTPException( @@ -65,6 +68,7 @@ async def get_current_user(token: str | None = Depends(oauth2_scheme), db: Sessi return user + # 获取已登录用户(抛出401如果未登录) async def get_required_user(user: User | None = Depends(get_current_user)): if user is None: @@ -75,6 +79,7 @@ async def get_required_user(user: User | None = Depends(get_current_user)): ) return user + # 获取管理员用户 async def get_admin_user(current_user: User = Depends(get_required_user)): if current_user.role not in ["admin", "superadmin"]: @@ -84,6 +89,7 @@ async def get_admin_user(current_user: User = Depends(get_required_user)): ) return current_user + # 获取超级管理员用户 async def get_superadmin_user(current_user: User = Depends(get_required_user)): if current_user.role != "superadmin": @@ -93,9 +99,10 @@ async def get_superadmin_user(current_user: User = Depends(get_required_user)): ) return current_user + # 检查路径是否为公开路径 def is_public_path(path: str) -> bool: - path = path.rstrip('/') # 去除尾部斜杠以便于匹配 + path = path.rstrip("/") # 去除尾部斜杠以便于匹配 for pattern in PUBLIC_PATHS: if re.match(pattern, path): return True diff --git a/server/utils/auth_utils.py b/server/utils/auth_utils.py index 4ea6fe98..08da0380 100644 --- a/server/utils/auth_utils.py +++ b/server/utils/auth_utils.py @@ -1,14 +1,16 @@ import hashlib import os -import jwt from datetime import datetime, timedelta from typing import Any +import jwt + # JWT配置 JWT_SECRET_KEY = os.environ.get("JWT_SECRET_KEY", "yuxi_know_secure_key") JWT_ALGORITHM = "HS256" JWT_EXPIRATION = 24 * 60 * 60 # 24小时过期 + class AuthUtils: """认证工具类""" diff --git a/server/utils/common_utils.py b/server/utils/common_utils.py index 73cce32c..1833f18b 100644 --- a/server/utils/common_utils.py +++ b/server/utils/common_utils.py @@ -1,19 +1,18 @@ """通用工具函数""" import logging -from sqlalchemy.orm import Session + from fastapi import Request -from server.models.user_model import User, OperationLog +from sqlalchemy.orm import Session + +from server.models.user_model import OperationLog, User def setup_logging(): """配置应用程序日志格式""" # 配置日志格式 logging.basicConfig( - level=logging.INFO, - format="%(asctime)s %(levelname)s: %(message)s", - datefmt="%Y-%m-%d %H:%M:%S", - force=True + level=logging.INFO, format="%(asctime)s %(levelname)s: %(message)s", datefmt="%Y-%m-%d %H:%M:%S", force=True ) # 确保uvicorn的日志也使用相同格式 @@ -21,10 +20,7 @@ def setup_logging(): uvicorn_access_logger = logging.getLogger("uvicorn.access") # 创建格式化器 - formatter = logging.Formatter( - fmt="%(asctime)s %(levelname)s: %(message)s", - datefmt="%m-%d %H:%M:%S" - ) + formatter = logging.Formatter(fmt="%(asctime)s %(levelname)s: %(message)s", datefmt="%m-%d %H:%M:%S") # 为所有处理器设置格式化器 for handler in uvicorn_logger.handlers: @@ -39,12 +35,7 @@ def log_operation(db: Session, user_id: int, operation: str, details: str = None if request: ip_address = request.client.host if request.client else None - log = OperationLog( - user_id=user_id, - operation=operation, - details=details, - ip_address=ip_address - ) + log = OperationLog(user_id=user_id, operation=operation, details=details, ip_address=ip_address) db.add(log) db.commit() @@ -60,6 +51,6 @@ def convert_serializable(obj): return [convert_serializable(item) for item in obj] if isinstance(obj, dict): return {k: convert_serializable(v) for k, v in obj.items()} - if hasattr(obj, '__dict__'): + if hasattr(obj, "__dict__"): return convert_serializable(vars(obj)) return obj diff --git a/src/__init__.py b/src/__init__.py index 49ecfe27..f9fa2cd7 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -1,38 +1,39 @@ import os + from dotenv import load_dotenv load_dotenv("src/.env", override=True) from concurrent.futures import ThreadPoolExecutor # noqa: E402 + executor = ThreadPoolExecutor() from src.config import Config # noqa: E402 + config = Config() # 导入知识库相关模块 +from src.knowledge.chroma_kb import ChromaKB # noqa: E402 from src.knowledge.kb_factory import KnowledgeBaseFactory # noqa: E402 from src.knowledge.kb_manager import KnowledgeBaseManager # noqa: E402 from src.knowledge.lightrag_kb import LightRagKB # noqa: E402 -from src.knowledge.chroma_kb import ChromaKB # noqa: E402 from src.knowledge.milvus_kb import MilvusKB # noqa: E402 # 注册知识库类型 -KnowledgeBaseFactory.register("chroma", ChromaKB, { - "chunk_size": 1000, - "chunk_overlap": 200, - "description": "基于 ChromaDB 的轻量级向量知识库,适合开发和小规模部署" -}) +KnowledgeBaseFactory.register( + "chroma", + ChromaKB, + {"chunk_size": 1000, "chunk_overlap": 200, "description": "基于 ChromaDB 的轻量级向量知识库,适合开发和小规模部署"}, +) -KnowledgeBaseFactory.register("milvus", MilvusKB, { - "chunk_size": 1000, - "chunk_overlap": 200, - "description": "基于 Milvus 的生产级向量知识库,适合大规模高性能部署" -}) +KnowledgeBaseFactory.register( + "milvus", + MilvusKB, + {"chunk_size": 1000, "chunk_overlap": 200, "description": "基于 Milvus 的生产级向量知识库,适合大规模高性能部署"}, +) -KnowledgeBaseFactory.register("lightrag", LightRagKB, { - "description": "基于图检索的知识库,支持实体关系构建和复杂查询" -}) +KnowledgeBaseFactory.register("lightrag", LightRagKB, {"description": "基于图检索的知识库,支持实体关系构建和复杂查询"}) # 创建知识库管理器 @@ -40,4 +41,5 @@ work_dir = os.path.join(config.save_dir, "knowledge_base_data") knowledge_base = KnowledgeBaseManager(work_dir) from src.knowledge import GraphDatabase # noqa: E402 + graph_base = GraphDatabase() diff --git a/src/agents/__init__.py b/src/agents/__init__.py index 3598fba7..18be435a 100644 --- a/src/agents/__init__.py +++ b/src/agents/__init__.py @@ -3,6 +3,7 @@ import asyncio from .chatbot.graph import ChatbotAgent from .react.graph import ReActAgent + class AgentManager: def __init__(self): self._classes = {} diff --git a/src/agents/chatbot/context.py b/src/agents/chatbot/context.py index 83abeb44..816fcc82 100644 --- a/src/agents/chatbot/context.py +++ b/src/agents/chatbot/context.py @@ -1,22 +1,18 @@ -from typing import Annotated from dataclasses import dataclass, field +from typing import Annotated from src.agents.common.context import BaseContext -from src.agents.common.tools import gen_tool_info from src.agents.common.mcp import MCP_SERVERS +from src.agents.common.tools import gen_tool_info from .tools import get_tools + @dataclass(kw_only=True) class Context(BaseContext): - model: Annotated[str, {"__template_metadata__": {"kind": "llm"}}] = field( default="siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507", - metadata={ - "name": "智能体模型", - "options": [], - "description": "智能体的驱动模型" - }, + metadata={"name": "智能体模型", "options": [], "description": "智能体的驱动模型"}, ) tools: Annotated[list[dict], {"__template_metadata__": {"kind": "tools"}}] = field( @@ -24,15 +20,11 @@ class Context(BaseContext): metadata={ "name": "工具", "options": gen_tool_info(get_tools()), # 这里的选择是所有的工具 - "description": "工具列表" + "description": "工具列表", }, ) mcps: list[str] = field( default_factory=list, - metadata={ - "name": "MCP服务器", - "options": list(MCP_SERVERS.keys()), - "description": "MCP服务器列表" - }, + metadata={"name": "MCP服务器", "options": list(MCP_SERVERS.keys()), "description": "MCP服务器列表"}, ) diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index b3e65187..5b138695 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -1,28 +1,26 @@ import os import uuid -from typing import Any, cast from pathlib import Path +from typing import Any, cast from langchain_core.messages import AIMessage, ToolMessage -from langgraph.graph import StateGraph, START, END -from langgraph.runtime import Runtime -from langgraph.prebuilt import ToolNode, tools_condition -from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite from langgraph.checkpoint.memory import InMemorySaver +from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite +from langgraph.graph import END, START, StateGraph +from langgraph.prebuilt import ToolNode, tools_condition +from langgraph.runtime import Runtime from src import config as sys_config -from src.utils import logger from src.agents.common.base import BaseAgent -from src.agents.common.models import load_chat_model from src.agents.common.mcp import get_mcp_tools +from src.agents.common.models import load_chat_model +from src.utils import logger - -from .state import State from .context import Context +from .state import State from .tools import get_tools - class ChatbotAgent(BaseAgent): name = "智能体助手" description = "基础的对话机器人,可以回答问题,默认不使用任何工具,可在配置中启用需要的工具。" @@ -69,16 +67,11 @@ class ChatbotAgent(BaseAgent): # 使用异步调用 response = cast( AIMessage, - await model.ainvoke( - [{"role": "system", "content": runtime.context.system_prompt}, *state.messages] - ), + await model.ainvoke([{"role": "system", "content": runtime.context.system_prompt}, *state.messages]), ) return {"messages": [response]} - - async def dynamic_tools_node( - self, state: State, runtime: Runtime[Context] - ) -> dict[str, list[ToolMessage]]: + async def dynamic_tools_node(self, state: State, runtime: Runtime[Context]) -> dict[str, list[ToolMessage]]: """Execute tools dynamically based on configuration. This function gets the available tools based on the current configuration @@ -133,6 +126,7 @@ class ChatbotAgent(BaseAgent): """获取异步存储实例""" return AsyncSqliteSaver(await self.get_async_conn()) + def main(): agent = ChatbotAgent(Context) @@ -140,6 +134,7 @@ def main(): config = {"configurable": {"thread_id": thread_id}} from src.agents.utils import agent_cli + agent_cli(agent, config) diff --git a/src/agents/chatbot/state.py b/src/agents/chatbot/state.py index c37963cf..28338901 100644 --- a/src/agents/chatbot/state.py +++ b/src/agents/chatbot/state.py @@ -2,12 +2,12 @@ from __future__ import annotations -from dataclasses import dataclass, field from collections.abc import Sequence +from dataclasses import dataclass, field +from typing import Annotated from langchain_core.messages import AnyMessage from langgraph.graph import add_messages -from typing import Annotated @dataclass @@ -17,6 +17,4 @@ class State: This class is used to define the initial state and structure of incoming data. """ - messages: Annotated[Sequence[AnyMessage], add_messages] = field( - default_factory=list - ) + messages: Annotated[Sequence[AnyMessage], add_messages] = field(default_factory=list) diff --git a/src/agents/chatbot/tools.py b/src/agents/chatbot/tools.py index 1d8fb52a..f4b921c2 100644 --- a/src/agents/chatbot/tools.py +++ b/src/agents/chatbot/tools.py @@ -2,8 +2,8 @@ from typing import Any from langchain_core.tools import tool -from src.utils import logger from src.agents.common.tools import get_buildin_tools +from src.utils import logger @tool @@ -26,9 +26,9 @@ def calculator(a: float, b: float, operation: str) -> float: logger.error(f"Calculator error: {e}") raise + def get_tools() -> dict[str, Any]: """获取所有可运行的工具(给大模型使用)""" tools = get_buildin_tools() tools.append(calculator) return tools - diff --git a/src/agents/common/base.py b/src/agents/common/base.py index 41a551f4..3d96cfb3 100644 --- a/src/agents/common/base.py +++ b/src/agents/common/base.py @@ -1,14 +1,14 @@ from __future__ import annotations from abc import abstractmethod + from langgraph.graph.state import CompiledStateGraph -from src.utils import logger from src.agents.common.context import BaseContext +from src.utils import logger class BaseAgent: - """ 定义一个基础 Agent 供 各类 graph 继承 """ @@ -23,7 +23,7 @@ class BaseAgent: @property def module_name(self) -> str: """Get the module name of the agent class.""" - return self.__class__.__module__.split('.')[-2] + return self.__class__.__module__.split(".")[-2] @property def id(self) -> str: @@ -42,18 +42,20 @@ class BaseAgent: async def get_config(self): return self.context_schema.from_file(module_name=self.module_name) - async def stream_values(self, messages: list[str], input_context = None, **kwargs): + async def stream_values(self, messages: list[str], input_context=None, **kwargs): graph = await self.get_graph() context = self.context_schema.from_file(module_name=self.module_name, input_context=input_context) for event in graph.astream({"messages": messages}, stream_mode="values", context=context): yield event["messages"] - async def stream_messages(self, messages: list[str], input_context = None, **kwargs): + async def stream_messages(self, messages: list[str], input_context=None, **kwargs): graph = await self.get_graph() context = self.context_schema.from_file(module_name=self.module_name, input_context=input_context) logger.debug(f"stream_messages: {context}") # TODO 的 Checkpointer 似乎还没有适配最新的 Context API - async for msg, metadata in graph.astream({"messages": messages}, stream_mode="messages", context=context, config={"configurable": input_context}): + async for msg, metadata in graph.astream( + {"messages": messages}, stream_mode="messages", context=context, config={"configurable": input_context} + ): yield msg, metadata async def check_checkpointer(self): @@ -76,12 +78,12 @@ class BaseAgent: result = [] if state: - messages = state.values.get('messages', []) + messages = state.values.get("messages", []) for msg in messages: - if hasattr(msg, 'model_dump'): + if hasattr(msg, "model_dump"): msg_dict = msg.model_dump() # 转换成字典 else: - msg_dict = dict(msg) if hasattr(msg, '__dict__') else {"content": str(msg)} + msg_dict = dict(msg) if hasattr(msg, "__dict__") else {"content": str(msg)} result.append(msg_dict) return result diff --git a/src/agents/common/context.py b/src/agents/common/context.py index e4dc1123..eb00f79e 100644 --- a/src/agents/common/context.py +++ b/src/agents/common/context.py @@ -3,11 +3,12 @@ from __future__ import annotations import os -import yaml import uuid -from dataclasses import dataclass, field, fields, MISSING +from dataclasses import MISSING, dataclass, field, fields from pathlib import Path -from typing import get_origin, get_args +from typing import get_args, get_origin + +import yaml from src import config as sys_config from src.utils import logger @@ -32,28 +33,17 @@ class BaseContext: thread_id: str = field( default_factory=lambda: str(uuid.uuid4()), - metadata={ - "name": "线程ID", - "configurable": False, - "description": "用来描述智能体的角色和行为" - }, + metadata={"name": "线程ID", "configurable": False, "description": "用来描述智能体的角色和行为"}, ) user_id: str = field( default_factory=lambda: str(uuid.uuid4()), - metadata={ - "name": "用户ID", - "configurable": False, - "description": "用来描述智能体的角色和行为" - }, + metadata={"name": "用户ID", "configurable": False, "description": "用来描述智能体的角色和行为"}, ) system_prompt: str = field( default="You are a helpful assistant.", - metadata={ - "name": "系统提示词", - "description": "用来描述智能体的角色和行为" - }, + metadata={"name": "系统提示词", "description": "用来描述智能体的角色和行为"}, ) @classmethod @@ -66,7 +56,7 @@ class BaseContext: if module_name is not None and os.path.exists(config_file_path): file_config = {} try: - with open(config_file_path, encoding='utf-8') as f: + with open(config_file_path, encoding="utf-8") as f: file_config = yaml.safe_load(f) or {} except Exception as e: logger.error(f"加载智能体配置文件出错: {e}") @@ -92,7 +82,7 @@ class BaseContext: config_file_path = Path(sys_config.save_dir) / "agents" / module_name / "config.yaml" # 确保目录存在 os.makedirs(os.path.dirname(config_file_path), exist_ok=True) - with open(config_file_path, 'w', encoding='utf-8') as f: + with open(config_file_path, "w", encoding="utf-8") as f: yaml.dump(configurable_config, f, indent=2, allow_unicode=True) return True @@ -118,7 +108,11 @@ class BaseContext: "type": type_name, "name": f.metadata.get("name", f.name), "options": f.metadata.get("options", []), - "default": f.default if f.default is not MISSING else f.default_factory() if f.default_factory is not MISSING else None, + "default": f.default + if f.default is not MISSING + else f.default_factory() + if f.default_factory is not MISSING + else None, "description": f.metadata.get("description", ""), "template_metadata": template_metadata, # Annotated 的额外元数据 } @@ -132,8 +126,8 @@ class BaseContext: if get_origin(field_type) is not None: # 处理泛型类型如 list[str], Annotated[str, {...}] origin = get_origin(field_type) - if hasattr(origin, '__name__'): - if origin.__name__ == 'Annotated': + if hasattr(origin, "__name__"): + if origin.__name__ == "Annotated": # Annotated 类型,获取真实类型 args = get_args(field_type) if args: @@ -141,7 +135,7 @@ class BaseContext: return origin.__name__ else: return str(origin) - elif hasattr(field_type, '__name__'): + elif hasattr(field_type, "__name__"): return field_type.__name__ else: return str(field_type) @@ -151,7 +145,7 @@ class BaseContext: """从 Annotated 类型中提取模板元数据""" if get_origin(field_type) is not None: origin = get_origin(field_type) - if hasattr(origin, '__name__') and origin.__name__ == 'Annotated': + if hasattr(origin, "__name__") and origin.__name__ == "Annotated": args = get_args(field_type) if len(args) > 1: # 查找包含 __template_metadata__ 的字典 diff --git a/src/agents/common/mcp.py b/src/agents/common/mcp.py index 9efe942e..ae0e5a05 100644 --- a/src/agents/common/mcp.py +++ b/src/agents/common/mcp.py @@ -1,13 +1,13 @@ """MCP Client setup and management for LangGraph ReAct Agent.""" import traceback -from typing import Any, cast from collections.abc import Callable +from typing import Any, cast -from langchain_mcp_adapters.tools import load_mcp_tools from langchain_mcp_adapters.client import ( # type: ignore[import-untyped] MultiServerMCPClient, ) +from langchain_mcp_adapters.tools import load_mcp_tools from src.utils import logger @@ -70,6 +70,7 @@ async def get_mcp_tools(server_name: str) -> list[Callable[..., Any]]: logger.opt(exception=True).warning(f"Failed to load tools from MCP server '{server_name}'") return [] + async def get_all_mcp_tools() -> list[Callable[..., Any]]: """Get all tools from all configured MCP servers.""" all_tools = [] diff --git a/src/agents/common/models.py b/src/agents/common/models.py index 7fe7fe02..befe1329 100644 --- a/src/agents/common/models.py +++ b/src/agents/common/models.py @@ -1,13 +1,12 @@ import os import traceback -from src import config -from src.utils import get_docker_safe_url -from src.models import get_custom_model from langchain_core.language_models import BaseChatModel from pydantic import SecretStr - +from src import config +from src.models import get_custom_model +from src.utils import get_docker_safe_url def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: @@ -18,6 +17,7 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: if provider == "custom": from langchain_openai import ChatOpenAI + model_info = get_custom_model(model) api_key = model_info.get("api_key") or "custom_model" base_url = get_docker_safe_url(model_info["api_base"]) @@ -34,6 +34,7 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: if provider in ["deepseek", "dashscope"]: from langchain_deepseek import ChatDeepSeek + return ChatDeepSeek( model=model, api_key=SecretStr(api_key), @@ -43,6 +44,7 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: elif provider == "together": from langchain_together import ChatTogether + return ChatTogether( model=model, api_key=SecretStr(api_key), @@ -52,6 +54,7 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: else: try: # 其他模型,默认使用OpenAIBase, like openai, zhipuai from langchain_openai import ChatOpenAI + return ChatOpenAI( model=model, api_key=SecretStr(api_key), diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index ef902058..b986fabe 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -1,10 +1,11 @@ -import inspect import asyncio +import inspect import traceback from typing import Annotated, Any -from pydantic import BaseModel, Field + from langchain_core.tools import StructuredTool, tool from langchain_tavily import TavilySearch +from pydantic import BaseModel, Field from src import config, graph_base, knowledge_base from src.utils import logger @@ -15,13 +16,16 @@ def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge """Use this to query knowledge graph, which include some food domain knowledge.""" try: logger.debug(f"Querying knowledge graph with: {query}") - result = graph_base.query_node(query, hops=2, return_format='triples') - logger.debug(f"Knowledge graph query returned {len(result.get('triples', [])) if isinstance(result, dict) else 'N/A'} triples") + result = graph_base.query_node(query, hops=2, return_format="triples") + logger.debug( + f"Knowledge graph query returned {len(result.get('triples', [])) if isinstance(result, dict) else 'N/A'} triples" + ) return result except Exception as e: logger.error(f"Knowledge graph query error: {e}, {traceback.format_exc()}") return f"知识图谱查询失败: {str(e)}" + def get_static_tools() -> list: """注册静态工具""" static_tools = [ @@ -38,13 +42,11 @@ def get_static_tools() -> list: class KnowledgeRetrieverModel(BaseModel): query_text: str = Field( description=( - "查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询," - "不要直接使用用户的原始输入去查询。" + "查询的关键词,查询的时候,应该尽量以可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。" ) ) - def get_kb_based_tools() -> list: """获取所有知识库基于的工具""" # 获取所有知识库 @@ -53,6 +55,7 @@ def get_kb_based_tools() -> list: def _create_retriever_wrapper(db_id: str, retriever_info: dict[str, Any]): """创建检索器包装函数的工厂函数,避免闭包变量捕获问题""" + async def async_retriever_wrapper(query_text: str) -> Any: """异步检索器包装函数""" retriever = retriever_info["retriever"] @@ -70,17 +73,13 @@ def get_kb_based_tools() -> list: return async_retriever_wrapper - for db_id, retrieve_info in retrievers.items(): try: # 使用改进的工具ID生成策略 tool_id = f"query_{db_id[:8]}" # 构建工具描述 - description = ( - f"使用 {retrieve_info['name']} 知识库进行检索。\n" - f"下面是这个知识库的描述:\n{retrieve_info['description'] or '没有描述。'} " - ) + description = f"使用 {retrieve_info['name']} 知识库进行检索。\n下面是这个知识库的描述:\n{retrieve_info['description'] or '没有描述。'} " # 使用工厂函数创建检索器包装函数,避免闭包问题 retriever_wrapper = _create_retriever_wrapper(db_id, retrieve_info) @@ -91,9 +90,7 @@ def get_kb_based_tools() -> list: name=tool_id, description=description, args_schema=KnowledgeRetrieverModel, - metadata=retrieve_info["metadata"] | { - "tag": ["knowledgebase"] - } + metadata=retrieve_info["metadata"] | {"tag": ["knowledgebase"]}, ) kb_tools.append(tool) @@ -129,30 +126,34 @@ def gen_tool_info(tools) -> list[dict[str, Any]]: # 获取注册的工具信息 for tool_obj in tools: try: - metadata = getattr(tool_obj, 'metadata', {}) or {} + metadata = getattr(tool_obj, "metadata", {}) or {} info = { "id": tool_obj.name, - "name": metadata.get('name', tool_obj.name), + "name": metadata.get("name", tool_obj.name), "description": tool_obj.description, - 'metadata': metadata, + "metadata": metadata, "args": [], # "is_async": is_async # Include async information } - if hasattr(tool_obj, 'args_schema') and tool_obj.args_schema: + if hasattr(tool_obj, "args_schema") and tool_obj.args_schema: schema = tool_obj.args_schema.schema() - for arg_name, arg_info in schema.get('properties', {}).items(): - info["args"].append({ - "name": arg_name, - "type": arg_info.get('type', ''), - "description": arg_info.get('description', '') - }) + for arg_name, arg_info in schema.get("properties", {}).items(): + info["args"].append( + { + "name": arg_name, + "type": arg_info.get("type", ""), + "description": arg_info.get("description", ""), + } + ) tools_info.append(info) # logger.debug(f"Successfully processed tool info for {tool_obj.name}") except Exception as e: - logger.error(f"Failed to process tool {getattr(tool_obj, 'name', 'unknown')}: {e}\n{traceback.format_exc()}") + logger.error( + f"Failed to process tool {getattr(tool_obj, 'name', 'unknown')}: {e}\n{traceback.format_exc()}" + ) continue except Exception as e: diff --git a/src/agents/common/utils.py b/src/agents/common/utils.py index ced1f75b..c6f70593 100644 --- a/src/agents/common/utils.py +++ b/src/agents/common/utils.py @@ -1,18 +1,17 @@ -from datetime import datetime, timezone, UTC import asyncio import os import traceback +from datetime import UTC, datetime, timezone -from src import config -from src.utils import get_docker_safe_url -from src.models import get_custom_model -from src.agents.common.base import BaseAgent from langchain_core.language_models import BaseChatModel -from langchain_core.runnables import RunnableConfig from langchain_core.messages import AIMessageChunk, ToolMessage +from langchain_core.runnables import RunnableConfig from pydantic import SecretStr - +from src import config +from src.agents.common.base import BaseAgent +from src.models import get_custom_model +from src.utils import get_docker_safe_url def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: @@ -23,6 +22,7 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: if provider == "custom": from langchain_openai import ChatOpenAI + model_info = get_custom_model(model) api_key = model_info.get("api_key") or "custom_model" base_url = get_docker_safe_url(model_info["api_base"]) @@ -39,6 +39,7 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: if provider in ["deepseek", "dashscope"]: from langchain_deepseek import ChatDeepSeek + return ChatDeepSeek( model=model, api_key=SecretStr(api_key), @@ -48,6 +49,7 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: elif provider == "together": from langchain_together import ChatTogether + return ChatTogether( model=model, api_key=SecretStr(api_key), @@ -57,6 +59,7 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: else: try: # 其他模型,默认使用OpenAIBase, like openai, zhipuai from langchain_openai import ChatOpenAI + return ChatOpenAI( model=model, api_key=SecretStr(api_key), @@ -66,7 +69,6 @@ def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel: raise ValueError(f"Model provider {provider} load failed, {e} \n {traceback.format_exc()}") - async def agent_cli(agent: BaseAgent, config: RunnableConfig | None = None): config = config or {} if "configurable" not in config: @@ -100,6 +102,6 @@ async def agent_cli(agent: BaseAgent, config: RunnableConfig | None = None): if isinstance(msg, ToolMessage): print(f"Tool: {msg.content}") + def get_cur_time_with_utc(): return datetime.now(tz=UTC).isoformat() - diff --git a/src/agents/react/graph.py b/src/agents/react/graph.py index ab7b966f..08e0a209 100644 --- a/src/agents/react/graph.py +++ b/src/agents/react/graph.py @@ -1,26 +1,27 @@ import os from pathlib import Path +from langchain_core.messages import AnyMessage, SystemMessage from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite from langgraph.prebuilt import create_react_agent -from langchain_core.messages import AnyMessage, SystemMessage from langgraph.runtime import get_runtime from src import config as sys_config -from src.utils import logger -from src.agents.common.context import BaseContext from src.agents.common.base import BaseAgent +from src.agents.common.context import BaseContext from src.agents.common.models import load_chat_model from src.agents.common.tools import get_buildin_tools - +from src.utils import logger model = load_chat_model("siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507") + def prompt(state) -> list[AnyMessage]: runtime = get_runtime(BaseContext) system_msg = SystemMessage(content=runtime.context.system_prompt) return [system_msg] + state["messages"] + class ReActAgent(BaseAgent): name = "ReAct (all tools)" description = "A react agent that can answer questions and help with tasks." @@ -38,12 +39,7 @@ class ReActAgent(BaseAgent): available_tools = get_buildin_tools() sqlite_checkpointer = AsyncSqliteSaver(await aiosqlite.connect(self.workdir / "react_history.db")) - graph = create_react_agent( - model, - tools=available_tools, - checkpointer=sqlite_checkpointer, - prompt=prompt - ) + graph = create_react_agent(model, tools=available_tools, checkpointer=sqlite_checkpointer, prompt=prompt) self.graph = graph logger.info("ReActAgent使用SQLite checkpointer构建成功") return graph diff --git a/src/config/__init__.py b/src/config/__init__.py index b1b82053..98f1eb6d 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -1,13 +1,15 @@ -import os import json -import yaml +import os from pathlib import Path + +import yaml + from src.utils.logging_config import logger -DEFAULT_MOCK_API = 'this_is_mock_api_key_in_frontend' +DEFAULT_MOCK_API = "this_is_mock_api_key_in_frontend" + class SimpleConfig(dict): - def __key(self, key): return "" if key is None else key # 目前忘记了这里为什么要 lower 了,只能说配置项最好不要有大写的 @@ -35,11 +37,10 @@ class SimpleConfig(dict): class Config(SimpleConfig): - def __init__(self): super().__init__() self._config_items = {} - self.save_dir = os.getenv('SAVE_DIR', 'saves') + self.save_dir = os.getenv("SAVE_DIR", "saves") self.filename = str(Path(f"{self.save_dir}/config/base.yaml")) os.makedirs(os.path.dirname(self.filename), exist_ok=True) @@ -48,7 +49,11 @@ class Config(SimpleConfig): ### >>> 默认配置 # 功能选项 self.add_item("enable_reranker", default=False, des="是否开启重排序") - self.add_item("enable_web_search", default=False, des="是否开启网页搜索(注:现阶段会根据 TAVILY_API_KEY 自动开启,无法手动配置,将会在下个版本移除此配置项)") # noqa: E501 + self.add_item( + "enable_web_search", + default=False, + des="是否开启网页搜索(注:现阶段会根据 TAVILY_API_KEY 自动开启,无法手动配置,将会在下个版本移除此配置项)", + ) # noqa: E501 # 默认智能体配置 self.add_item("default_agent_id", default="", des="默认智能体ID") # 模型配置 @@ -57,8 +62,18 @@ class Config(SimpleConfig): self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(self.model_names.keys())) self.add_item("model_name", default="zai-org/GLM-4.5", des="模型名称") - self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(self.embed_model_names.keys())) - self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(self.reranker_names.keys())) # noqa: E501 + self.add_item( + "embed_model", + default="siliconflow/BAAI/bge-m3", + des="Embedding 模型", + choices=list(self.embed_model_names.keys()), + ) + self.add_item( + "reranker", + default="siliconflow/BAAI/bge-reranker-v2-m3", + des="Re-Ranker 模型", + choices=list(self.reranker_names.keys()), + ) # noqa: E501 ### <<< 默认配置结束 self.load() @@ -66,11 +81,7 @@ class Config(SimpleConfig): def add_item(self, key, default, des=None, choices=None): self.__setattr__(key, default) - self._config_items[key] = { - "default": default, - "des": des, - "choices": choices - } + self._config_items[key] = {"default": default, "des": des, "choices": choices} def __dict__(self): blocklist = [ @@ -87,12 +98,12 @@ class Config(SimpleConfig): 从 models.yaml 和 models.private.yml 中更新 MODEL_NAMES """ - with open(Path("src/static/models.yaml"), encoding='utf-8') as f: + with open(Path("src/static/models.yaml"), encoding="utf-8") as f: _models = yaml.safe_load(f) # 尝试打开一个 models.private.yml 文件,用来覆盖 models.yaml 中的配置 try: - with open(Path("src/static/models.private.yml"), encoding='utf-8') as f: + with open(Path("src/static/models.private.yml"), encoding="utf-8") as f: _models_private = yaml.safe_load(f) except FileNotFoundError: _models_private = {} @@ -110,7 +121,7 @@ class Config(SimpleConfig): "EMBED_MODEL_INFO": self.embed_model_names, "RERANKER_LIST": self.reranker_names, } - with open(Path("src/static/models.private.yml"), 'w', encoding='utf-8') as f: + with open(Path("src/static/models.private.yml"), "w", encoding="utf-8") as f: yaml.dump(_models, f, indent=2, allow_unicode=True) def handle_self(self): @@ -121,10 +132,14 @@ class Config(SimpleConfig): if self.model_dir: if os.path.exists(self.model_dir): - logger.debug(f"The model directory ({self.model_dir}) contains the following folders: {os.listdir(self.model_dir)}") + logger.debug( + f"The model directory ({self.model_dir}) contains the following folders: {os.listdir(self.model_dir)}" + ) else: - logger.warning(f"Warning: The model directory ({self.model_dir}) does not exist. If not configured, please ignore it. If configured, please check if the configuration is correct;" - "For example, the mapping in the docker-compose file") + logger.warning( + f"Warning: The model directory ({self.model_dir}) does not exist. If not configured, please ignore it. If configured, please check if the configuration is correct;" + "For example, the mapping in the docker-compose file" + ) # 检查模型提供商的环境变量 conds = {} @@ -138,13 +153,14 @@ class Config(SimpleConfig): self.enable_web_search = True self.valuable_model_provider = [k for k, v in self.model_provider_status.items() if v] - assert len(self.valuable_model_provider) > 0, f"No model provider available, please check your `.env` file. API_KEY_LIST: {conds}" + assert len(self.valuable_model_provider) > 0, ( + f"No model provider available, please check your `.env` file. API_KEY_LIST: {conds}" + ) def load(self): """根据传入的文件覆盖掉默认配置""" logger.info(f"Loading config from {self.filename}") if self.filename is not None and os.path.exists(self.filename): - if self.filename.endswith(".json"): with open(self.filename) as f: content = f.read() @@ -173,14 +189,14 @@ class Config(SimpleConfig): os.makedirs(os.path.dirname(self.filename), exist_ok=True) if self.filename.endswith(".json"): - with open(self.filename, 'w+') as f: + with open(self.filename, "w+") as f: json.dump(self.__dict__(), f, indent=4, ensure_ascii=False) elif self.filename.endswith(".yaml"): - with open(self.filename, 'w+') as f: + with open(self.filename, "w+") as f: yaml.dump(self.__dict__(), f, indent=2, allow_unicode=True) else: logger.warning(f"Unknown config file type {self.filename}, save as json") - with open(self.filename, 'w+') as f: + with open(self.filename, "w+") as f: json.dump(self, f, indent=4) logger.info(f"Config file {self.filename} saved") diff --git a/src/knowledge/chroma_kb.py b/src/knowledge/chroma_kb.py index 00613188..17b1a806 100644 --- a/src/knowledge/chroma_kb.py +++ b/src/knowledge/chroma_kb.py @@ -1,21 +1,26 @@ +import json import os import time import traceback -import json -from pathlib import Path -from typing import Optional, Any from datetime import datetime +from pathlib import Path +from typing import Any, Optional import chromadb +from chromadb.api.types import Documents, EmbeddingFunction, Embeddings from chromadb.config import Settings -from chromadb.api.types import EmbeddingFunction, Documents, Embeddings from chromadb.utils.embedding_functions import OpenAIEmbeddingFunction -from src.knowledge.indexing import process_url_to_markdown, process_file_to_markdown -from src.knowledge.knowledge_base import KnowledgeBase -from src.knowledge.kb_utils import split_text_into_chunks, split_text_into_qa_chunks, prepare_item_metadata, get_embedding_config -from src.utils import logger, hashstr from src import config +from src.knowledge.indexing import process_file_to_markdown, process_url_to_markdown +from src.knowledge.kb_utils import ( + get_embedding_config, + prepare_item_metadata, + split_text_into_chunks, + split_text_into_qa_chunks, +) +from src.knowledge.knowledge_base import KnowledgeBase +from src.utils import hashstr, logger class ChromaKB(KnowledgeBase): @@ -40,8 +45,7 @@ class ChromaKB(KnowledgeBase): # 初始化 ChromaDB 客户端 self.chroma_client = chromadb.PersistentClient( - path=self.chroma_db_path, - settings=Settings(anonymized_telemetry=False) + path=self.chroma_db_path, settings=Settings(anonymized_telemetry=False) ) # 存储集合映射 {db_id: collection} @@ -68,10 +72,7 @@ class ChromaKB(KnowledgeBase): try: # 尝试获取现有集合 - collection = self.chroma_client.get_collection( - name=collection_name, - embedding_function=embedding_function - ) + collection = self.chroma_client.get_collection(name=collection_name, embedding_function=embedding_function) logger.info(f"Retrieved existing collection: {collection_name}") # 检查现有集合的配置是否匹配当前的 embed_info @@ -82,7 +83,9 @@ class ChromaKB(KnowledgeBase): logger.debug(f"Collection {collection_name} uses model '{current_model}', but expected '{expected_model}'.") # 如果模型不匹配,删除现有集合并重新创建 if current_model != expected_model: - logger.warning(f"Collection {collection_name} uses model '{current_model}', but expected '{expected_model}'. Recreating collection.") + logger.warning( + f"Collection {collection_name} uses model '{current_model}', but expected '{expected_model}'. Recreating collection." + ) self.chroma_client.delete_collection(name=collection_name) raise Exception("Model mismatch, recreating collection") @@ -92,12 +95,10 @@ class ChromaKB(KnowledgeBase): collection_metadata = { "db_id": db_id, "created_at": datetime.now().isoformat(), - "embedding_model": embed_info.get("name") if embed_info else "default" + "embedding_model": embed_info.get("name") if embed_info else "default", } collection = self.chroma_client.create_collection( - name=collection_name, - embedding_function=embedding_function, - metadata=collection_metadata + name=collection_name, embedding_function=embedding_function, metadata=collection_metadata ) logger.info(f"Created new collection: {collection_name}") @@ -114,7 +115,7 @@ class ChromaKB(KnowledgeBase): return OpenAIEmbeddingFunction( model_name=config_dict["model"], api_key=config_dict["api_key"], - api_base=config_dict["base_url"].replace('/embeddings', '') + api_base=config_dict["base_url"].replace("/embeddings", ""), ) async def _get_chroma_collection(self, db_id: str): @@ -141,11 +142,11 @@ class ChromaKB(KnowledgeBase): def _split_text_into_chunks(self, text: str, file_id: str, filename: str, params: dict) -> list[dict]: """将文本分割成块""" # 检查是否使用QA分割模式 - use_qa_split = params.get('use_qa_split', False) + use_qa_split = params.get("use_qa_split", False) if use_qa_split: # 使用QA分割模式 - qa_separator = params.get('qa_separator', '\n\n\n') + qa_separator = params.get("qa_separator", "\n\n\n") chunks = split_text_into_qa_chunks(text, file_id, filename, qa_separator, params) else: # 使用传统分割模式 @@ -157,13 +158,12 @@ class ChromaKB(KnowledgeBase): "source": chunk["source"], "chunk_id": chunk["chunk_id"], "full_doc_id": file_id, - "chunk_type": chunk.get("chunk_type", "normal") # 添加chunk类型标识 + "chunk_type": chunk.get("chunk_type", "normal"), # 添加chunk类型标识 } return chunks - async def add_content(self, db_id: str, items: list[str], - params:dict | None) -> list[dict]: + async def add_content(self, db_id: str, items: list[str], params: dict | None) -> list[dict]: """添加内容(文件/URL)""" if db_id not in self.databases_meta: raise ValueError(f"Database {db_id} not found") @@ -172,7 +172,7 @@ class ChromaKB(KnowledgeBase): if not collection: raise ValueError(f"Failed to get ChromaDB collection for {db_id}") - content_type = params.get('content_type', 'file') if params else 'file' + content_type = params.get("content_type", "file") if params else "file" processed_items_info = [] for item in items: @@ -205,24 +205,20 @@ class ChromaKB(KnowledgeBase): ids = [chunk["id"] for chunk in chunks] # 插入到 ChromaDB - collection.add( - documents=documents, - metadatas=metadatas, - ids=ids - ) + collection.add(documents=documents, metadatas=metadatas, ids=ids) logger.info(f"Inserted {content_type} {item} into ChromaDB. Done.") # 更新状态为完成 self.files_meta[file_id]["status"] = "done" self._save_metadata() - file_record['status'] = "done" + file_record["status"] = "done" except Exception as e: logger.error(f"处理{content_type} {item} 失败: {e}, {traceback.format_exc()}") self.files_meta[file_id]["status"] = "failed" self._save_metadata() - file_record['status'] = "failed" + file_record["status"] = "failed" finally: self._remove_from_processing_queue(file_id) @@ -241,9 +237,7 @@ class ChromaKB(KnowledgeBase): similarity_threshold = kwargs.get("similarity_threshold", 0.0) results = collection.query( - query_texts=[query_text], - n_results=top_k, - include=["documents", "metadatas", "distances"] + query_texts=[query_text], n_results=top_k, include=["documents", "metadatas", "distances"] ) if not results or not results.get("documents") or not results["documents"][0]: @@ -262,14 +256,10 @@ class ChromaKB(KnowledgeBase): metadata = metadatas[i] if i < len(metadatas) else {} # 确保 file_id 在元数据中,并使用统一的键名 - if 'full_doc_id' in metadata: - metadata['file_id'] = metadata.pop('full_doc_id') + if "full_doc_id" in metadata: + metadata["file_id"] = metadata.pop("full_doc_id") - retrieved_chunks.append({ - "content": doc, - "metadata": metadata, - "score": similarity - }) + retrieved_chunks.append({"content": doc, "metadata": metadata, "score": similarity}) logger.debug(f"ChromaDB query response: {len(retrieved_chunks)} chunks found (after similarity filtering)") return retrieved_chunks @@ -284,10 +274,7 @@ class ChromaKB(KnowledgeBase): if collection: try: # 查找所有相关的chunks - results = collection.get( - where={"full_doc_id": file_id}, - include=["metadatas"] - ) + results = collection.get(where={"full_doc_id": file_id}, include=["metadatas"]) # 删除所有相关chunks if results and results.get("ids"): @@ -312,10 +299,7 @@ class ChromaKB(KnowledgeBase): if collection: try: # 获取文档的所有chunks - results = collection.get( - where={"full_doc_id": file_id}, - include=["documents", "metadatas"] - ) + results = collection.get(where={"full_doc_id": file_id}, include=["documents", "metadatas"]) # 构建chunks数据 doc_chunks = [] @@ -325,7 +309,9 @@ class ChromaKB(KnowledgeBase): "id": chunk_id, "content": results["documents"][i] if i < len(results["documents"]) else "", "metadata": results["metadatas"][i] if i < len(results["metadatas"]) else {}, - "chunk_order_index": results["metadatas"][i].get("chunk_index", i) if i < len(results["metadatas"]) else i + "chunk_order_index": results["metadatas"][i].get("chunk_index", i) + if i < len(results["metadatas"]) + else i, } doc_chunks.append(chunk_data) diff --git a/src/knowledge/graphbase.py b/src/knowledge/graphbase.py index 823bd15c..b80edbb2 100644 --- a/src/knowledge/graphbase.py +++ b/src/knowledge/graphbase.py @@ -1,7 +1,7 @@ -import os import json -import warnings +import os import traceback +import warnings from neo4j import GraphDatabase as GD from neo4j import Query @@ -15,6 +15,7 @@ warnings.filterwarnings("ignore", category=UserWarning) UIE_MODEL = None + class GraphDatabase: def __init__(self): self.driver = None @@ -55,10 +56,11 @@ class GraphDatabase: """检查图数据库是否正在运行""" return self.status == "open" or self.status == "processing" - def get_sample_nodes(self, kgdb_name='neo4j', num=50): + def get_sample_nodes(self, kgdb_name="neo4j", num=50): """获取指定数据库的 num 个节点信息,优先返回连通的节点子图""" assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) + def query(tx, num): """Note: 使用连通性查询获取集中的节点子图""" # 首先尝试获取一个连通的子图 @@ -104,25 +106,25 @@ class GraphDatabase: try: results = tx.run(query_str, num=int(num)) - formatted_results = {'nodes': [], 'edges': []} + formatted_results = {"nodes": [], "edges": []} node_ids = set() for item in results: - h_node = item['h'] - t_node = item['t'] + h_node = item["h"] + t_node = item["t"] # 避免重复添加节点 - if h_node['id'] not in node_ids: - formatted_results['nodes'].append(h_node) - node_ids.add(h_node['id']) - if t_node['id'] not in node_ids: - formatted_results['nodes'].append(t_node) - node_ids.add(t_node['id']) + if h_node["id"] not in node_ids: + formatted_results["nodes"].append(h_node) + node_ids.add(h_node["id"]) + if t_node["id"] not in node_ids: + formatted_results["nodes"].append(t_node) + node_ids.add(t_node["id"]) - formatted_results['edges'].append(item['r']) + formatted_results["edges"].append(item["r"]) # 如果连通查询没有返回足够的结果,回退到原始查询 - if len(formatted_results['nodes']) < num // 2: + if len(formatted_results["nodes"]) < num // 2: fallback_query = """ MATCH (n:Entity)-[r]-(m:Entity) WHERE elementId(n) < elementId(m) @@ -133,11 +135,11 @@ class GraphDatabase: LIMIT $num """ fallback_results = tx.run(fallback_query, num=int(num)) - formatted_results = {'nodes': [], 'edges': []} + formatted_results = {"nodes": [], "edges": []} for item in fallback_results: - formatted_results['nodes'].extend([item['h'], item['t']]) - formatted_results['edges'].append(item['r']) + formatted_results["nodes"].extend([item["h"], item["t"]]) + formatted_results["edges"].append(item["r"]) return formatted_results @@ -154,11 +156,11 @@ class GraphDatabase: LIMIT $num """ results = tx.run(fallback_query, num=int(num)) - formatted_results = {'nodes': [], 'edges': []} + formatted_results = {"nodes": [], "edges": []} for item in results: - formatted_results['nodes'].extend([item['h'], item['t']]) - formatted_results['edges'].append(item['r']) + formatted_results["nodes"].extend([item["h"], item["t"]]) + formatted_results["edges"].append(item["r"]) return formatted_results @@ -171,7 +173,7 @@ class GraphDatabase: assert self.driver is not None, "Database is not connected" with self.driver.session() as session: existing_databases = session.run("SHOW DATABASES") - existing_db_names = [db['name'] for db in existing_databases] + existing_db_names = [db["name"] for db in existing_databases] if existing_db_names: print(f"已存在数据库: {existing_db_names[0]}") @@ -183,15 +185,17 @@ class GraphDatabase: def use_database(self, kgdb_name="neo4j"): """切换到指定数据库""" - assert kgdb_name == self.kgdb_name, f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致" + assert kgdb_name == self.kgdb_name, ( + f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致" + ) if self.status == "closed": self.start() - - async def txt_add_vector_entity(self, triples, kgdb_name='neo4j'): + async def txt_add_vector_entity(self, triples, kgdb_name="neo4j"): """添加实体三元组""" assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) + def _index_exists(tx, index_name): """检查索引是否存在""" result = tx.run("SHOW INDEXES") @@ -203,11 +207,16 @@ class GraphDatabase: def _create_graph(tx, data): """添加一个三元组""" for entry in data: - tx.run(""" + tx.run( + """ MERGE (h:Entity:Upload {name: $h}) MERGE (t:Entity:Upload {name: $t}) MERGE (h)-[r:RELATION {type: $r}]->(t) - """, h=entry['h'], t=entry['t'], r=entry['r']) + """, + h=entry["h"], + t=entry["t"], + r=entry["r"], + ) def _create_vector_index(tx, dim): """创建向量索引""" @@ -232,42 +241,50 @@ class GraphDatabase: param_placeholders = ", ".join([f"${key}" for key in params.keys()]) # 执行查询 - result = tx.run(f""" + result = tx.run( + f""" MATCH (n:Entity) WHERE n.name IN [{param_placeholders}] AND n.embedding IS NULL RETURN n.name AS name - """, params) + """, + params, + ) return [record["name"] for record in result] def _batch_set_embeddings(tx, entity_embedding_pairs): """批量设置实体的嵌入向量""" for entity_name, embedding in entity_embedding_pairs: - tx.run(""" + tx.run( + """ MATCH (e:Entity {name: $name}) CALL db.create.setNodeVectorProperty(e, 'embedding', $embedding) - """, name=entity_name, embedding=embedding) + """, + name=entity_name, + embedding=embedding, + ) # 判断模型名称是否匹配 self.embed_model_name = self.embed_model_name or config.embed_model cur_embed_info = config.embed_model_names.get(self.embed_model_name) logger.warning(f"embed_model_name={self.embed_model_name}, {cur_embed_info=}") - assert self.embed_model_name == config.embed_model or self.embed_model_name is None, \ + assert self.embed_model_name == config.embed_model or self.embed_model_name is None, ( f"embed_model_name={self.embed_model_name}, {config.embed_model=}" + ) with self.driver.session() as session: logger.info(f"Adding entity to {kgdb_name}") session.execute_write(_create_graph, triples) logger.info(f"Creating vector index for {kgdb_name} with {config.embed_model}") - session.execute_write(_create_vector_index, cur_embed_info['dimension']) + session.execute_write(_create_vector_index, cur_embed_info["dimension"]) # 收集所有需要处理的实体名称,去重 all_entities = [] for entry in triples: - if entry['h'] not in all_entities: - all_entities.append(entry['h']) - if entry['t'] not in all_entities: - all_entities.append(entry['t']) + if entry["h"] not in all_entities: + all_entities.append(entry["h"]) + if entry["t"] not in all_entities: + all_entities.append(entry["t"]) # 筛选出没有embedding的节点 nodes_without_embedding = session.execute_read(_get_nodes_without_embedding, all_entities) @@ -282,11 +299,9 @@ class GraphDatabase: total_entities = len(nodes_without_embedding) for i in range(0, total_entities, max_batch_size): - batch_entities = nodes_without_embedding[i:i+max_batch_size] + batch_entities = nodes_without_embedding[i : i + max_batch_size] logger.debug( - f"Processing entities batch " - f"{i//max_batch_size + 1}/{(total_entities-1)//max_batch_size + 1} " - f"({len(batch_entities)} entities)" + f"Processing entities batch {i // max_batch_size + 1}/{(total_entities - 1) // max_batch_size + 1} ({len(batch_entities)} entities)" ) # 批量获取嵌入向量 @@ -301,15 +316,15 @@ class GraphDatabase: # 数据添加完成后保存图信息 self.save_graph_info() - async def jsonl_file_add_entity(self, file_path, kgdb_name='neo4j'): + async def jsonl_file_add_entity(self, file_path, kgdb_name="neo4j"): assert self.driver is not None, "Database is not connected" self.status = "processing" - kgdb_name = kgdb_name or 'neo4j' + kgdb_name = kgdb_name or "neo4j" self.use_database(kgdb_name) # 切换到指定数据库 logger.info(f"Start adding entity to {kgdb_name} with {file_path}") def read_triples(file_path): - with open(file_path, encoding='utf-8') as file: + with open(file_path, encoding="utf-8") as file: for line in file: if line.strip(): yield json.loads(line.strip()) @@ -347,7 +362,9 @@ class GraphDatabase: """ tx.run(query) - def query_node(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, max_entities=5, return_format='graph', **kwargs): + def query_node( + self, entity_name, threshold=0.9, kgdb_name="neo4j", hops=2, max_entities=5, return_format="graph", **kwargs + ): """知识图谱查询节点的入口:""" assert self.driver is not None, "Database is not connected" assert self.is_running(), "图数据库未启动" @@ -363,29 +380,33 @@ class GraphDatabase: logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}") # 对每个合格的实体进行查询 - all_query_results = {'nodes': [], 'edges': [], 'triples': []} + all_query_results = {"nodes": [], "edges": [], "triples": []} for entity in qualified_entities: query_result = self._query_specific_entity(entity_name=entity, kgdb_name=kgdb_name, hops=hops) - if return_format == 'graph': - all_query_results['nodes'].extend(query_result['nodes']) - all_query_results['edges'].extend(query_result['edges']) - elif return_format == 'triples': - all_query_results['triples'].extend(query_result['triples']) + if return_format == "graph": + all_query_results["nodes"].extend(query_result["nodes"]) + all_query_results["edges"].extend(query_result["edges"]) + elif return_format == "triples": + all_query_results["triples"].extend(query_result["triples"]) else: raise ValueError(f"Invalid return_format: {return_format}") return all_query_results - def _query_with_fuzzy_match(self, keyword, kgdb_name='neo4j'): + def _query_with_fuzzy_match(self, keyword, kgdb_name="neo4j"): """模糊查询""" assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) + def query_fuzzy_match(tx, keyword): - result = tx.run(""" + result = tx.run( + """ MATCH (n:Entity) WHERE n.name CONTAINS $keyword RETURN DISTINCT n.name AS name - """, keyword=keyword) + """, + keyword=keyword, + ) values = result.values() logger.debug(f"Fuzzy Query Results: {values}") return values @@ -393,7 +414,7 @@ class GraphDatabase: with self.driver.session() as session: return session.execute_read(query_fuzzy_match, keyword) - def _query_with_vector_sim(self, keyword, kgdb_name='neo4j', threshold=0.9): + def _query_with_vector_sim(self, keyword, kgdb_name="neo4j", threshold=0.9): """向量查询""" assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) @@ -409,22 +430,26 @@ class GraphDatabase: def query_by_vector(tx, text, threshold): # 首先检查索引是否存在 if not _index_exists(tx, "entityEmbeddings"): - raise Exception("向量索引不存在,请先创建索引,或当前图谱中未上传任何三元组(知识库中自动构建的,不会在此处展示和检索)。") + raise Exception( + "向量索引不存在,请先创建索引,或当前图谱中未上传任何三元组(知识库中自动构建的,不会在此处展示和检索)。" + ) embedding = self.get_embedding(text) - result = tx.run(""" + result = tx.run( + """ CALL db.index.vector.queryNodes('entityEmbeddings', 10, $embedding) YIELD node AS similarEntity, score RETURN similarEntity.name AS name, score - """, embedding=embedding) + """, + embedding=embedding, + ) return [r for r in result if r["score"] > threshold] with self.driver.session() as session: results = session.execute_read(query_by_vector, keyword, threshold=threshold) return results - - def _query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2, limit=100): + def _query_specific_entity(self, entity_name, kgdb_name="neo4j", hops=2, limit=100): """查询指定实体三元组信息(无向关系)""" assert self.driver is not None, "Database is not connected" if not entity_name: @@ -467,12 +492,12 @@ class GraphDatabase: logger.info(f"未找到实体 {entity_name} 的相关信息") return {} - formatted_results = {'nodes': [], 'edges': [], 'triples': []} + formatted_results = {"nodes": [], "edges": [], "triples": []} for item in results: - formatted_results['nodes'].extend([item['h'], item['t']]) - formatted_results['edges'].append(item['r']) - formatted_results['triples'].append((item['h']['name'], item['r']['type'], item['t']['name'])) + formatted_results["nodes"].extend([item["h"], item["t"]]) + formatted_results["edges"].append(item["r"]) + formatted_results["triples"].append((item["h"]["name"], item["r"]["type"], item["t"]["name"])) logger.debug(f"Query Results: {results}") return formatted_results @@ -506,20 +531,27 @@ class GraphDatabase: return outputs def set_embedding(self, tx, entity_name, embedding): - tx.run(""" + tx.run( + """ MATCH (e:Entity {name: $name}) CALL db.create.setNodeVectorProperty(e, 'embedding', $embedding) - """, name=entity_name, embedding=embedding) + """, + name=entity_name, + embedding=embedding, + ) def get_graph_info(self, graph_name="neo4j"): assert self.driver is not None, "Database is not connected" self.use_database(graph_name) + def query(tx): # 只统计包含Entity标签的节点 entity_count = tx.run("MATCH (n:Entity) RETURN count(n) AS count").single()["count"] # 只统计包含RELATION标签的关系 relationship_count = tx.run("MATCH ()-[r:RELATION]->() RETURN count(r) AS count").single()["count"] - triples_count = tx.run("MATCH (n:Entity)-[r:RELATION]->(m:Entity) RETURN count(n) AS count").single()["count"] + triples_count = tx.run("MATCH (n:Entity)-[r:RELATION]->(m:Entity) RETURN count(n) AS count").single()[ + "count" + ] # 获取所有标签 labels = tx.run("CALL db.labels() YIELD label RETURN collect(label) AS labels").single()["labels"] @@ -532,7 +564,7 @@ class GraphDatabase: "labels": labels, "status": self.status, "embed_model_name": self.embed_model_name, - "unindexed_node_count": self.query_nodes_without_embedding(graph_name) + "unindexed_node_count": self.query_nodes_without_embedding(graph_name), } try: @@ -543,6 +575,7 @@ class GraphDatabase: # 添加时间戳 from datetime import datetime + graph_info["last_updated"] = datetime.now().isoformat() return graph_info else: @@ -565,7 +598,7 @@ class GraphDatabase: return False info_file_path = os.path.join(self.work_dir, "graph_info.json") - with open(info_file_path, 'w', encoding='utf-8') as f: + with open(info_file_path, "w", encoding="utf-8") as f: json.dump(graph_info, f, ensure_ascii=False, indent=2) # logger.info(f"图数据库信息已保存到:{info_file_path}") @@ -574,7 +607,7 @@ class GraphDatabase: logger.error(f"保存图数据库信息失败:{e}") return False - def query_nodes_without_embedding(self, kgdb_name='neo4j'): + def query_nodes_without_embedding(self, kgdb_name="neo4j"): """查询没有嵌入向量的节点 Returns: @@ -605,7 +638,7 @@ class GraphDatabase: logger.debug(f"图数据库信息文件不存在:{info_file_path}") return False - with open(info_file_path, encoding='utf-8') as f: + with open(info_file_path, encoding="utf-8") as f: graph_info = json.load(f) # 更新对象属性 @@ -621,7 +654,7 @@ class GraphDatabase: logger.error(f"加载图数据库信息失败:{e}") return False - def add_embedding_to_nodes(self, node_names=None, kgdb_name='neo4j'): + def add_embedding_to_nodes(self, node_names=None, kgdb_name="neo4j"): """为节点添加嵌入向量 Args: @@ -655,13 +688,12 @@ class GraphDatabase: edges = [] for item in results: - nodes.extend([item['h'], item['t']]) - edges.append(item['r']) + nodes.extend([item["h"], item["t"]]) + edges.append(item["r"]) formatted_results = {"nodes": nodes, "edges": edges} return formatted_results - def _extract_relationship_info(self, relationship, source_name=None, target_name=None, node_dict=None): """ 提取关系信息并返回格式化的节点和边信息 @@ -701,12 +733,13 @@ class GraphDatabase: return node_info, edge_info + def clean_triples_embedding(triples): for item in triples: - if hasattr(item[0], '_properties'): - item[0]._properties['embedding'] = None - if hasattr(item[2], '_properties'): - item[2]._properties['embedding'] = None + if hasattr(item[0], "_properties"): + item[0]._properties["embedding"] = None + if hasattr(item[2], "_properties"): + item[2]._properties["embedding"] = None return triples diff --git a/src/knowledge/indexing.py b/src/knowledge/indexing.py index b6590a29..c758e74f 100644 --- a/src/knowledge/indexing.py +++ b/src/knowledge/indexing.py @@ -1,16 +1,17 @@ -import os import asyncio +import os from pathlib import Path + from langchain.schema.document import Document from langchain.text_splitter import RecursiveCharacterTextSplitter from langchain_community.document_loaders import ( - TextLoader, - PyPDFLoader, - Docx2txtLoader, - UnstructuredMarkdownLoader, - UnstructuredHTMLLoader, CSVLoader, - JSONLoader + Docx2txtLoader, + JSONLoader, + PyPDFLoader, + TextLoader, + UnstructuredHTMLLoader, + UnstructuredMarkdownLoader, ) from src.utils import hashstr, logger @@ -31,22 +32,22 @@ def chunk_with_parser(file_path, params=None): file_type = Path(file_path).suffix.lower() # 选择合适的加载器 - if file_type in ['.txt']: + if file_type in [".txt"]: loader = TextLoader(file_path) - elif file_type in ['.md']: + elif file_type in [".md"]: loader = UnstructuredMarkdownLoader(file_path) - elif file_type in ['.docx', '.doc']: + elif file_type in [".docx", ".doc"]: loader = Docx2txtLoader(file_path) - elif file_type in ['.html', '.htm']: + elif file_type in [".html", ".htm"]: loader = UnstructuredHTMLLoader(file_path) - elif file_type in ['.json']: + elif file_type in [".json"]: loader = JSONLoader(file_path, jq_schema=".") - elif file_type in ['.csv']: + elif file_type in [".csv"]: loader = CSVLoader(file_path) else: @@ -73,6 +74,7 @@ def chunk_with_parser(file_path, params=None): return nodes + def chunk_text(text, params=None): """ 将文本切分成固定大小的块 @@ -83,9 +85,7 @@ def chunk_text(text, params=None): # 创建文本分割器 text_splitter = RecursiveCharacterTextSplitter( - chunk_size=chunk_size, - chunk_overlap=chunk_overlap, - separators=["\n\n", "\n", ".", " ", ""] + chunk_size=chunk_size, chunk_overlap=chunk_overlap, separators=["\n\n", "\n", ".", " ", ""] ) # 分割文档 @@ -95,9 +95,11 @@ def chunk_text(text, params=None): nodes = [{"text": node, "metadata": {"chunk_idx": i}} for i, node in enumerate(nodes)] return nodes + def chunk(text_or_path, params=None): raise NotImplementedError("chunk is deprecated, use chunk_with_parser or chunk_text instead") + def pdfreader(file_path, params=None): """读取PDF文件并返回text文本""" if isinstance(file_path, str): @@ -114,6 +116,7 @@ def pdfreader(file_path, params=None): text = "\n\n".join([d.page_content for d in docs]) return text + def plainreader(file_path): """读取普通文本文件并返回text文本""" assert os.path.exists(file_path), "File not found" @@ -124,6 +127,7 @@ def plainreader(file_path): text = "\n\n".join([d.page_content for d in docs]) return text + def parse_pdf(file, params=None): """ 解析PDF文件,支持多种OCR方式 @@ -149,14 +153,17 @@ def parse_pdf(file, params=None): try: if opt_ocr == "onnx_rapid_ocr": from src.plugins import ocr + return ocr.process_pdf(file, params=params) elif opt_ocr == "mineru_ocr": from src.plugins import ocr + return ocr.process_file_mineru(file, params=params) elif opt_ocr == "paddlex_ocr": from src.plugins import ocr + return ocr.process_file_paddlex(file, params=params) else: @@ -167,11 +174,7 @@ def parse_pdf(file, params=None): raise except Exception as e: logger.error(f"PDF parsing failed: {str(e)}") - raise OCRServiceException( - f"PDF解析失败: {str(e)}", - opt_ocr, - "parsing_failed" - ) + raise OCRServiceException(f"PDF解析失败: {str(e)}", opt_ocr, "parsing_failed") def parse_image(file, params=None): @@ -190,14 +193,17 @@ def parse_image(file, params=None): try: if opt_ocr == "onnx_rapid_ocr": from src.plugins import ocr + return ocr.process_image(file, params=params) elif opt_ocr == "mineru_ocr": from src.plugins import ocr + return ocr.process_file_mineru(file, params=params) elif opt_ocr == "paddlex_ocr": from src.plugins import ocr + return ocr.process_file_paddlex(file, params=params) else: @@ -208,18 +214,17 @@ def parse_image(file, params=None): raise except Exception as e: logger.error(f"Image parsing failed: {str(e)}") - raise OCRServiceException( - f"Image解析失败: {str(e)}", - opt_ocr, - "parsing_failed" - ) + raise OCRServiceException(f"Image解析失败: {str(e)}", opt_ocr, "parsing_failed") + async def parse_pdf_async(file, params=None): return await asyncio.to_thread(parse_pdf, file, params=params) + async def parse_image_async(file, params=None): return await asyncio.to_thread(parse_image, file, params=params) + async def process_file_to_markdown(file_path: str, params: dict | None = None) -> str: """ 将不同类型的文件转换为markdown格式 @@ -234,40 +239,43 @@ async def process_file_to_markdown(file_path: str, params: dict | None = None) - file_path_obj = Path(file_path) file_ext = file_path_obj.suffix.lower() - if file_ext == '.pdf': + if file_ext == ".pdf": # 使用 OCR 处理 PDF text = await parse_pdf_async(str(file_path_obj), params=params) return f"# {file_path_obj.name}\n\n{text}" - elif file_ext in ['.txt', '.md']: + elif file_ext in [".txt", ".md"]: # 直接读取文本文件 - with open(file_path_obj, encoding='utf-8') as f: + with open(file_path_obj, encoding="utf-8") as f: content = f.read() return f"# {file_path_obj.name}\n\n{content}" - elif file_ext in ['.doc', '.docx']: + elif file_ext in [".doc", ".docx"]: # 处理 Word 文档 from docx import Document # type: ignore + doc = Document(file_path_obj) - text = '\n'.join([para.text for para in doc.paragraphs]) + text = "\n".join([para.text for para in doc.paragraphs]) return f"# {file_path_obj.name}\n\n{text}" - elif file_ext in ['.jpg', '.jpeg', '.png', '.bmp', '.tiff', '.tif']: + elif file_ext in [".jpg", ".jpeg", ".png", ".bmp", ".tiff", ".tif"]: # 使用 OCR 处理图片 text = await parse_image_async(str(file_path_obj), params=params) return f"# {file_path_obj.name}\n\n{text}" - elif file_ext in ['.html', '.htm']: + elif file_ext in [".html", ".htm"]: # 使用 BeautifulSoup 处理 HTML 文件 from markdownify import markdownify as md - with open(file_path_obj, encoding='utf-8') as f: + + with open(file_path_obj, encoding="utf-8") as f: content = f.read() text = md(content, heading_style="ATX") return f"# {file_path_obj.name}\n\n{text}" - elif file_ext == '.csv': + elif file_ext == ".csv": # 处理 CSV 文件 import pandas as pd + df = pd.read_csv(file_path_obj) # 将每一行数据与表头组合成独立的表格 markdown_content = f"# {file_path_obj.name}\n\n" @@ -280,9 +288,10 @@ async def process_file_to_markdown(file_path: str, params: dict | None = None) - return markdown_content.strip() - elif file_ext in ['.xls', '.xlsx']: + elif file_ext in [".xls", ".xlsx"]: # 处理 Excel 文件 import pandas as pd + # 读取所有工作表 excel_file = pd.ExcelFile(file_path_obj) markdown_content = f"# {file_path_obj.name}\n\n" @@ -300,10 +309,11 @@ async def process_file_to_markdown(file_path: str, params: dict | None = None) - return markdown_content.strip() - elif file_ext == '.json': + elif file_ext == ".json": # 处理 JSON 文件 import json - with open(file_path_obj, encoding='utf-8') as f: + + with open(file_path_obj, encoding="utf-8") as f: data = json.load(f) # 将 JSON 数据格式化为 markdown 代码块 json_str = json.dumps(data, ensure_ascii=False, indent=2) @@ -313,6 +323,7 @@ async def process_file_to_markdown(file_path: str, params: dict | None = None) - # 尝试作为文本文件读取 raise ValueError(f"Unsupported file type: {file_ext}") + async def process_url_to_markdown(url: str, params: dict | None = None) -> str: """ 将URL转换为markdown格式 @@ -329,10 +340,9 @@ async def process_url_to_markdown(url: str, params: dict | None = None) -> str: try: response = requests.get(url, timeout=30) - soup = BeautifulSoup(response.content, 'html.parser') + soup = BeautifulSoup(response.content, "html.parser") text_content = soup.get_text() return f"# {url}\n\n{text_content}" except Exception as e: logger.error(f"Failed to process URL {url}: {e}") return f"# {url}\n\nFailed to process URL: {e}" - diff --git a/src/knowledge/kb_factory.py b/src/knowledge/kb_factory.py index b1e4dbc1..35549ca3 100644 --- a/src/knowledge/kb_factory.py +++ b/src/knowledge/kb_factory.py @@ -1,5 +1,6 @@ from typing import Any -from src.knowledge.knowledge_base import KnowledgeBase, KBNotFoundError + +from src.knowledge.knowledge_base import KBNotFoundError, KnowledgeBase from src.utils import logger @@ -13,8 +14,7 @@ class KnowledgeBaseFactory: _default_configs: dict[str, dict] = {} @classmethod - def register(cls, kb_type: str, kb_class: type[KnowledgeBase], - default_config: dict = None): + def register(cls, kb_type: str, kb_class: type[KnowledgeBase], default_config: dict = None): """ 注册知识库类型 @@ -48,8 +48,7 @@ class KnowledgeBaseFactory: """ if kb_type not in cls._kb_types: available_types = list(cls._kb_types.keys()) - raise KBNotFoundError(f"Unknown knowledge base type: {kb_type}. " - f"Available types: {available_types}") + raise KBNotFoundError(f"Unknown knowledge base type: {kb_type}. Available types: {available_types}") kb_class = cls._kb_types[kb_type] @@ -79,7 +78,7 @@ class KnowledgeBaseFactory: result[kb_type] = { "class_name": kb_class.__name__, "description": kb_class.__doc__ or "", - "default_config": cls._default_configs[kb_type] + "default_config": cls._default_configs[kb_type], } return result diff --git a/src/knowledge/kb_manager.py b/src/knowledge/kb_manager.py index eea9f7e1..de750d8a 100644 --- a/src/knowledge/kb_manager.py +++ b/src/knowledge/kb_manager.py @@ -1,12 +1,12 @@ -import os -import json -import time import asyncio -from typing import Any +import json +import os +import time from datetime import datetime +from typing import Any -from src.knowledge.knowledge_base import KnowledgeBase, KBNotFoundError, KBOperationError from src.knowledge.kb_factory import KnowledgeBaseFactory +from src.knowledge.knowledge_base import KBNotFoundError, KBOperationError, KnowledgeBase from src.utils import logger @@ -49,7 +49,7 @@ class KnowledgeBaseManager: meta_file = os.path.join(self.work_dir, "global_metadata.json") if os.path.exists(meta_file): try: - with open(meta_file, encoding='utf-8') as f: + with open(meta_file, encoding="utf-8") as f: data = json.load(f) self.global_databases_meta = data.get("databases", {}) logger.info(f"Loaded global metadata for {len(self.global_databases_meta)} databases") @@ -59,12 +59,8 @@ class KnowledgeBaseManager: def _save_global_metadata(self): """保存全局元数据""" meta_file = os.path.join(self.work_dir, "global_metadata.json") - data = { - "databases": self.global_databases_meta, - "updated_at": datetime.now().isoformat(), - "version": "2.0" - } - with open(meta_file, 'w', encoding='utf-8') as f: + data = {"databases": self.global_databases_meta, "updated_at": datetime.now().isoformat(), "version": "2.0"} + with open(meta_file, "w", encoding="utf-8") as f: json.dump(data, f, ensure_ascii=False, indent=2) def _initialize_existing_kbs(self): @@ -140,7 +136,9 @@ class KnowledgeBaseManager: return {"databases": all_databases} - async def create_database(self, database_name: str, description: str, kb_type: str, embed_info: dict | None = None, **kwargs) -> dict: + async def create_database( + self, database_name: str, description: str, kb_type: str, embed_info: dict | None = None, **kwargs + ) -> dict: """ 创建数据库 @@ -160,8 +158,7 @@ class KnowledgeBaseManager: kb_instance = self._get_or_create_kb_instance(kb_type) - db_info = kb_instance.create_database(database_name, description, - embed_info, **kwargs) + db_info = kb_instance.create_database(database_name, description, embed_info, **kwargs) db_id = db_info["db_id"] async with self._metadata_lock: @@ -170,7 +167,7 @@ class KnowledgeBaseManager: "description": description, "kb_type": kb_type, "created_at": datetime.now().isoformat(), - "additional_params": kwargs.copy() + "additional_params": kwargs.copy(), } self._save_global_metadata() @@ -203,7 +200,7 @@ class KnowledgeBaseManager: kb_instance = self._get_kb_for_database(db_id) return await kb_instance.aquery(query_text, db_id, **kwargs) - async def export_data(self, db_id: str, format: str = 'zip', **kwargs) -> str: + async def export_data(self, db_id: str, format: str = "zip", **kwargs) -> str: """导出知识库数据""" kb_instance = self._get_kb_for_database(db_id) return await kb_instance.export_data(db_id, format=format, **kwargs) @@ -294,17 +291,13 @@ class KnowledgeBaseManager: info[kb_type] = { "work_dir": kb_instance.work_dir, "database_count": len(kb_instance.databases_meta), - "file_count": len(kb_instance.files_meta) + "file_count": len(kb_instance.files_meta), } return info def get_statistics(self) -> dict: """获取统计信息""" - stats = { - "total_databases": len(self.global_databases_meta), - "kb_types": {}, - "total_files": 0 - } + stats = {"total_databases": len(self.global_databases_meta), "kb_types": {}, "total_files": 0} # 按知识库类型统计 for db_meta in self.global_databases_meta.values(): @@ -352,7 +345,7 @@ class KnowledgeBaseManager: kb_instance = self._get_kb_for_database(db_id) # 如果不是 LightRagKB 实例,返回错误 - if not hasattr(kb_instance, '_get_lightrag_instance'): + if not hasattr(kb_instance, "_get_lightrag_instance"): logger.error(f"Knowledge base instance for {db_id} is not LightRagKB") return None diff --git a/src/knowledge/kb_utils.py b/src/knowledge/kb_utils.py index cce8d899..c94449d2 100644 --- a/src/knowledge/kb_utils.py +++ b/src/knowledge/kb_utils.py @@ -2,9 +2,11 @@ import os import time from pathlib import Path from typing import Any + from langchain_text_splitters import MarkdownTextSplitter -from src.utils import hashstr, get_docker_safe_url, logger + from src import config +from src.utils import get_docker_safe_url, hashstr, logger def split_text_into_chunks(text: str, file_id: str, filename: str, params: dict = {}) -> list[dict]: @@ -12,8 +14,8 @@ def split_text_into_chunks(text: str, file_id: str, filename: str, params: dict 将文本分割成块,使用 LangChain 的 MarkdownTextSplitter 进行智能分割 """ chunks = [] - chunk_size = params.get('chunk_size', 1000) - chunk_overlap = params.get('chunk_overlap', 200) + chunk_size = params.get("chunk_size", 1000) + chunk_overlap = params.get("chunk_overlap", 200) # 使用 MarkdownTextSplitter 进行智能分割 # MarkdownTextSplitter 会尝试沿着 Markdown 格式的标题进行分割 @@ -27,15 +29,17 @@ def split_text_into_chunks(text: str, file_id: str, filename: str, params: dict # 转换为标准格式 for chunk_index, chunk_content in enumerate(text_chunks): if chunk_content.strip(): # 跳过空块 - chunks.append({ - "id": f"{file_id}_chunk_{chunk_index}", - "content": chunk_content.strip(), - "file_id": file_id, - "filename": filename, - "chunk_index": chunk_index, - "source": filename, - "chunk_id": f"{file_id}_chunk_{chunk_index}" - }) + chunks.append( + { + "id": f"{file_id}_chunk_{chunk_index}", + "content": chunk_content.strip(), + "file_id": file_id, + "filename": filename, + "chunk_index": chunk_index, + "source": filename, + "chunk_id": f"{file_id}_chunk_{chunk_index}", + } + ) logger.debug(f"Successfully split text into {len(chunks)} chunks using MarkdownTextSplitter") return chunks @@ -64,15 +68,16 @@ def prepare_item_metadata(item: str, content_type: str, db_id: str) -> dict: "file_type": file_type, "status": "processing", "created_at": time.time(), - "file_id": file_id + "file_id": file_id, } -def split_text_into_qa_chunks(text: str, file_id: str, filename: str, - qa_separator: None | str = None, params: dict = {}) -> list[dict]: +def split_text_into_qa_chunks( + text: str, file_id: str, filename: str, qa_separator: None | str = None, params: dict = {} +) -> list[dict]: """ 将文本按QA对分割成块,使用 LangChain 的 CharacterTextSplitter 进行分割""" - qa_separator = qa_separator or '\n\n' + qa_separator = qa_separator or "\n\n" text_chunks = text.split(qa_separator) # 转换为标准格式 @@ -80,19 +85,23 @@ def split_text_into_qa_chunks(text: str, file_id: str, filename: str, for chunk_index, chunk_content in enumerate(text_chunks): if chunk_content.strip(): # 跳过空块 chunk_content = chunk_content.strip()[:4096] - chunks.append({ - "id": f"{file_id}_qa_chunk_{chunk_index}", - "content": chunk_content.strip(), - "file_id": file_id, - "filename": filename, - "chunk_index": chunk_index, - "source": filename, - "chunk_id": f"{file_id}_qa_chunk_{chunk_index}", - "chunk_type": "qa" # 标识为QA类型的chunk - }) + chunks.append( + { + "id": f"{file_id}_qa_chunk_{chunk_index}", + "content": chunk_content.strip(), + "file_id": file_id, + "filename": filename, + "chunk_index": chunk_index, + "source": filename, + "chunk_id": f"{file_id}_qa_chunk_{chunk_index}", + "chunk_type": "qa", # 标识为QA类型的chunk + } + ) logger.debug(f"QA chunks: {chunks[0]}") - logger.debug(f"Successfully split QA text into {len(chunks)} chunks using CharacterTextSplitter with `{qa_separator=}`") + logger.debug( + f"Successfully split QA text into {len(chunks)} chunks using CharacterTextSplitter with `{qa_separator=}`" + ) return chunks @@ -110,17 +119,18 @@ def get_embedding_config(embed_info: dict) -> dict: try: if embed_info: - config_dict['model'] = embed_info["name"] - config_dict['api_key'] = os.getenv(embed_info["api_key"], embed_info["api_key"]) - config_dict['base_url'] = embed_info["base_url"] - config_dict['dimension'] = embed_info.get("dimension", 1024) + config_dict["model"] = embed_info["name"] + config_dict["api_key"] = os.getenv(embed_info["api_key"], embed_info["api_key"]) + config_dict["base_url"] = embed_info["base_url"] + config_dict["dimension"] = embed_info.get("dimension", 1024) else: from src.models import select_embedding_model + default_model = select_embedding_model(config.embed_model) - config_dict['model'] = default_model.model - config_dict['api_key'] = default_model.api_key - config_dict['base_url'] = default_model.base_url - config_dict['dimension'] = getattr(default_model, 'dimension', 1024) + config_dict["model"] = default_model.model + config_dict["api_key"] = default_model.api_key + config_dict["base_url"] = default_model.base_url + config_dict["dimension"] = getattr(default_model, "dimension", 1024) except Exception as e: logger.error(f"Error in get_embedding_config: {e}, {embed_info}") diff --git a/src/knowledge/knowledge_base.py b/src/knowledge/knowledge_base.py index 5516aa63..7b67cb06 100644 --- a/src/knowledge/knowledge_base.py +++ b/src/knowledge/knowledge_base.py @@ -1,27 +1,30 @@ -import os import json +import os import time from abc import ABC, abstractmethod -from typing import Any from collections.abc import AsyncGenerator -from pathlib import Path from datetime import datetime +from pathlib import Path +from typing import Any from src.utils import logger class KnowledgeBaseException(Exception): """知识库统一异常基类""" + pass class KBNotFoundError(KnowledgeBaseException): """知识库不存在错误""" + pass class KBOperationError(KnowledgeBaseException): """知识库操作错误""" + pass @@ -84,8 +87,14 @@ class KnowledgeBase(ABC): """ pass - def create_database(self, database_name: str, description: str, - embed_info: dict | None = None, llm_info: dict | None = None, **kwargs) -> dict: + def create_database( + self, + database_name: str, + description: str, + embed_info: dict | None = None, + llm_info: dict | None = None, + **kwargs, + ) -> dict: """ 创建数据库 @@ -110,7 +119,7 @@ class KnowledgeBase(ABC): "embed_info": embed_info, "llm_info": llm_info, "metadata": kwargs, - "created_at": datetime.now().isoformat() + "created_at": datetime.now().isoformat(), } self._save_metadata() @@ -137,8 +146,7 @@ class KnowledgeBase(ABC): """ if db_id in self.databases_meta: # 删除相关文件记录 - files_to_delete = [fid for fid, finfo in self.files_meta.items() - if finfo.get("database_id") == db_id] + files_to_delete = [fid for fid, finfo in self.files_meta.items() if finfo.get("database_id") == db_id] for file_id in files_to_delete: del self.files_meta[file_id] @@ -150,6 +158,7 @@ class KnowledgeBase(ABC): working_dir = os.path.join(self.work_dir, db_id) if os.path.exists(working_dir): import shutil + try: shutil.rmtree(working_dir) except Exception as e: @@ -158,8 +167,7 @@ class KnowledgeBase(ABC): return {"message": "删除成功"} @abstractmethod - async def add_content(self, db_id: str, items: list[str], - params: dict | None = None) -> list[dict]: + async def add_content(self, db_id: str, items: list[str], params: dict | None = None) -> list[dict]: """ 添加内容(文件/URL) @@ -188,7 +196,7 @@ class KnowledgeBase(ABC): """ pass - async def export_data(self, db_id: str, format: str = 'zip', **kwargs) -> str: + async def export_data(self, db_id: str, format: str = "zip", **kwargs) -> str: pass def query(self, query_text: str, db_id: str, **kwargs) -> list[dict]: @@ -204,6 +212,7 @@ class KnowledgeBase(ABC): 一个包含字典的列表,每个字典代表一个检索到的文档块。 """ import asyncio + logger.warning("query is deprecated, use aquery instead") return asyncio.run(self.aquery(query_text, db_id, **kwargs)) @@ -236,7 +245,7 @@ class KnowledgeBase(ABC): "path": file_info.get("path", ""), "type": file_info.get("file_type", ""), "status": file_info.get("status", "done"), - "created_at": file_info.get("created_at", time.time()) + "created_at": file_info.get("created_at", time.time()), } # 按创建时间倒序排序文件列表 @@ -272,7 +281,7 @@ class KnowledgeBase(ABC): "path": file_info.get("path", ""), "type": file_info.get("file_type", ""), "status": file_info.get("status", "done"), - "created_at": file_info.get("created_at", time.time()) + "created_at": file_info.get("created_at", time.time()), } # 按创建时间倒序排序文件列表 @@ -323,8 +332,6 @@ class KnowledgeBase(ABC): with cls._processing_lock: return file_id in cls._processing_files - - def _check_and_fix_processing_status(self, db_id: str) -> None: """ 检查并修复异常的processing状态 @@ -338,14 +345,16 @@ class KnowledgeBase(ABC): # 检查该数据库下所有processing状态的文件 for file_id, file_info in self.files_meta.items(): - if (file_info.get("database_id") == db_id and - file_info.get("status") == "processing"): - + if file_info.get("database_id") == db_id and file_info.get("status") == "processing": # 检查文件是否真的在处理队列中 if not self._is_file_in_processing_queue(file_id): - logger.warning(f"File {file_id} has processing status but is not in processing queue, marking as error") + logger.warning( + f"File {file_id} has processing status but is not in processing queue, marking as error" + ) self.files_meta[file_id]["status"] = "error" - self.files_meta[file_id]["error"] = "Processing interrupted - file not found in processing queue" + self.files_meta[file_id]["error"] = ( + "Processing interrupted - file not found in processing queue" + ) status_changed = True # 如果有状态变更,保存元数据 @@ -430,16 +439,18 @@ class KnowledgeBase(ABC): """ retrievers = {} for db_id, meta in self.databases_meta.items(): + def make_retriever(db_id): async def retriever(query_text): return await self.aquery(query_text, db_id) + return retriever retrievers[db_id] = { "name": meta["name"], "description": meta["description"], "retriever": make_retriever(db_id), - "metadata": meta + "metadata": meta, } return retrievers @@ -448,7 +459,7 @@ class KnowledgeBase(ABC): meta_file = os.path.join(self.work_dir, f"metadata_{self.kb_type}.json") if os.path.exists(meta_file): try: - with open(meta_file, encoding='utf-8') as f: + with open(meta_file, encoding="utf-8") as f: data = json.load(f) self.databases_meta = data.get("databases", {}) self.files_meta = data.get("files", {}) @@ -464,9 +475,9 @@ class KnowledgeBase(ABC): "databases": self.databases_meta, "files": self.files_meta, "kb_type": self.kb_type, - "updated_at": datetime.now().isoformat() + "updated_at": datetime.now().isoformat(), } - with open(meta_file, 'w', encoding='utf-8') as f: + with open(meta_file, "w", encoding="utf-8") as f: json.dump(data, f, ensure_ascii=False, indent=2) except Exception as e: logger.error(f"Failed to save {self.kb_type} metadata: {e}") diff --git a/src/knowledge/lightrag_kb.py b/src/knowledge/lightrag_kb.py index 31675043..c8c0ec36 100644 --- a/src/knowledge/lightrag_kb.py +++ b/src/knowledge/lightrag_kb.py @@ -3,17 +3,16 @@ import traceback from datetime import datetime from lightrag import LightRAG, QueryParam +from lightrag.kg.shared_storage import initialize_pipeline_status from lightrag.llm.openai import openai_complete_if_cache, openai_embed from lightrag.utils import EmbeddingFunc, setup_logger -from lightrag.kg.shared_storage import initialize_pipeline_status -from pymilvus import connections, utility from neo4j import GraphDatabase +from pymilvus import connections, utility +from src.knowledge.indexing import process_file_to_markdown, process_url_to_markdown +from src.knowledge.kb_utils import get_embedding_config, prepare_item_metadata from src.knowledge.knowledge_base import KnowledgeBase -from src.knowledge.indexing import process_url_to_markdown, process_file_to_markdown -from src.knowledge.kb_utils import prepare_item_metadata, get_embedding_config -from src.utils import logger, hashstr - +from src.utils import hashstr, logger LIGHTRAG_LLM_PROVIDER = os.getenv("LIGHTRAG_LLM_PROVIDER", "siliconflow") LIGHTRAG_LLM_NAME = os.getenv("LIGHTRAG_LLM_NAME", "zai-org/GLM-4.5-Air") @@ -38,8 +37,9 @@ class LightRagKB(KnowledgeBase): # 设置 LightRAG 日志 log_dir = os.path.join(work_dir, "logs", "lightrag") os.makedirs(log_dir, exist_ok=True) - setup_logger("lightrag", log_file_path=os.path.join( - log_dir, f"lightrag_{datetime.now().strftime('%Y-%m-%d')}.log")) + setup_logger( + "lightrag", log_file_path=os.path.join(log_dir, f"lightrag_{datetime.now().strftime('%Y-%m-%d')}.log") + ) logger.info("LightRagKB initialized") @@ -52,15 +52,11 @@ class LightRagKB(KnowledgeBase): """删除数据库,同时清除Milvus和Neo4j中的数据""" # Drop Milvus collection try: - milvus_uri = os.getenv('MILVUS_URI', 'http://localhost:19530') - milvus_token = os.getenv('MILVUS_TOKEN', '') + milvus_uri = os.getenv("MILVUS_URI", "http://localhost:19530") + milvus_token = os.getenv("MILVUS_TOKEN", "") connection_alias = f"lightrag_{hashstr(db_id, 6)}" - connections.connect( - alias=connection_alias, - uri=milvus_uri, - token=milvus_token - ) + connections.connect(alias=connection_alias, uri=milvus_uri, token=milvus_token) # 删除 LightRAG 创建的三个集合 collection_names = [f"{db_id}_chunks", f"{db_id}_relationships", f"{db_id}_entities"] @@ -76,24 +72,28 @@ class LightRagKB(KnowledgeBase): logger.error(f"Failed to drop Milvus collection {db_id}: {e}") # Delete Neo4j data - neo4j_uri = os.getenv('NEO4J_URI', 'bolt://localhost:7687') - neo4j_username = os.getenv('NEO4J_USERNAME', 'neo4j') - neo4j_password = os.getenv('NEO4J_PASSWORD', '0123456789') + neo4j_uri = os.getenv("NEO4J_URI", "bolt://localhost:7687") + neo4j_username = os.getenv("NEO4J_USERNAME", "neo4j") + neo4j_password = os.getenv("NEO4J_PASSWORD", "0123456789") try: driver = GraphDatabase.driver(neo4j_uri, auth=(neo4j_username, neo4j_password)) with driver.session() as session: # 删除带有特定 db_id 标签的节点和关系 - session.run(""" - MATCH (n:`""" + db_id + """`) + session.run( + """ + MATCH (n:`""" + + db_id + + """`) DETACH DELETE n - """) + """ + ) logger.info(f"Deleted Neo4j nodes and relationships for workspace {db_id}") except Exception as e: logger.error(f"Failed to delete Neo4j data for {db_id}: {e}") finally: - if 'driver' in locals(): + if "driver" in locals(): driver.close() # Delete local files and metadata @@ -194,6 +194,7 @@ class LightRagKB(KnowledgeBase): base_url=model.base_url, **kwargs, ) + return llm_model_func def _get_embedding_func(self, embed_info: dict): @@ -211,8 +212,7 @@ class LightRagKB(KnowledgeBase): ), ) - async def add_content(self, db_id: str, items: list[str], - params: dict | None = None) -> list[dict]: + async def add_content(self, db_id: str, items: list[str], params: dict | None = None) -> list[dict]: """添加内容(文件/URL)""" if db_id not in self.databases_meta: raise ValueError(f"Database {db_id} not found") @@ -221,7 +221,7 @@ class LightRagKB(KnowledgeBase): if not rag: raise ValueError(f"Failed to get LightRAG instance for {db_id}") - content_type = params.get('content_type', 'file') if params else 'file' + content_type = params.get("content_type", "file") if params else "file" processed_items_info = [] for item in items: @@ -240,24 +240,20 @@ class LightRagKB(KnowledgeBase): # 根据内容类型处理内容 if content_type == "file": markdown_content = await process_file_to_markdown(item, params=params) - markdown_content_lines = markdown_content[:100].replace('\n', ' ') + markdown_content_lines = markdown_content[:100].replace("\n", " ") logger.info(f"Markdown content: {markdown_content_lines}...") else: # URL markdown_content = await process_url_to_markdown(item, params=params) # 使用 LightRAG 插入内容 - await rag.ainsert( - input=markdown_content, - ids=file_id, - file_paths=item_path - ) + await rag.ainsert(input=markdown_content, ids=file_id, file_paths=item_path) logger.info(f"Inserted {content_type} {item} into LightRAG. Done.") # 更新状态为完成 self.files_meta[file_id]["status"] = "done" self._save_metadata() - file_record['status'] = "done" + file_record["status"] = "done" except Exception as e: error_msg = str(e) @@ -265,8 +261,8 @@ class LightRagKB(KnowledgeBase): self.files_meta[file_id]["status"] = "failed" self.files_meta[file_id]["error"] = error_msg self._save_metadata() - file_record['status'] = "failed" - file_record['error'] = error_msg + file_record["status"] = "failed" + file_record["error"] = error_msg finally: self._remove_from_processing_queue(file_id) @@ -324,7 +320,7 @@ class LightRagKB(KnowledgeBase): if rag: try: # 获取文档的所有 chunks - assert hasattr(rag.text_chunks, 'get_all'), "text_chunks does not have get_all method" + assert hasattr(rag.text_chunks, "get_all"), "text_chunks does not have get_all method" all_chunks = await rag.text_chunks.get_all() # type: ignore # 筛选属于该文档的 chunks @@ -344,7 +340,7 @@ class LightRagKB(KnowledgeBase): return {"lines": []} - async def export_data(self, db_id: str, format: str = 'csv', **kwargs) -> str: + async def export_data(self, db_id: str, format: str = "csv", **kwargs) -> str: """ 使用 LightRAG 原生功能导出知识库数据。 [注意] 此功能当前已禁用。 diff --git a/src/knowledge/milvus_kb.py b/src/knowledge/milvus_kb.py index fe1cba60..a758f359 100644 --- a/src/knowledge/milvus_kb.py +++ b/src/knowledge/milvus_kb.py @@ -1,24 +1,26 @@ +import asyncio +import json import os import time import traceback -import json -import asyncio -from pathlib import Path -from typing import Any from datetime import datetime from functools import partial +from pathlib import Path +from typing import Any -from pymilvus import ( - connections, utility, Collection, CollectionSchema, - FieldSchema, DataType, db -) +from pymilvus import Collection, CollectionSchema, DataType, FieldSchema, connections, db, utility from src import config -from src.models.embedding import OtherEmbedding +from src.knowledge.indexing import process_file_to_markdown, process_url_to_markdown +from src.knowledge.kb_utils import ( + get_embedding_config, + prepare_item_metadata, + split_text_into_chunks, + split_text_into_qa_chunks, +) from src.knowledge.knowledge_base import KnowledgeBase -from src.knowledge.indexing import process_url_to_markdown, process_file_to_markdown -from src.knowledge.kb_utils import split_text_into_chunks, split_text_into_qa_chunks, prepare_item_metadata, get_embedding_config -from src.utils import logger, hashstr +from src.models.embedding import OtherEmbedding +from src.utils import hashstr, logger MILVUS_AVAILABLE = True @@ -42,9 +44,9 @@ class MilvusKB(KnowledgeBase): # Milvus 配置 # self.milvus_host = kwargs.get('milvus_host', os.getenv('MILVUS_HOST', 'localhost')) # self.milvus_port = kwargs.get('milvus_port', int(os.getenv('MILVUS_PORT', '19530'))) - self.milvus_token = kwargs.get('milvus_token', os.getenv('MILVUS_TOKEN', '')) - self.milvus_uri = kwargs.get('milvus_uri', os.getenv('MILVUS_URI', 'http://localhost:19530')) - self.milvus_db = kwargs.get('milvus_db', 'yuxi_know') + self.milvus_token = kwargs.get("milvus_token", os.getenv("MILVUS_TOKEN", "")) + self.milvus_uri = kwargs.get("milvus_uri", os.getenv("MILVUS_URI", "http://localhost:19530")) + self.milvus_db = kwargs.get("milvus_db", "yuxi_know") # 连接名称 self.connection_alias = f"milvus_{hashstr(work_dir, 6)}" @@ -53,8 +55,8 @@ class MilvusKB(KnowledgeBase): self.collections: dict[str, Any] = {} # 分块配置 - self.chunk_size = kwargs.get('chunk_size', 1000) - self.chunk_overlap = kwargs.get('chunk_overlap', 200) + self.chunk_size = kwargs.get("chunk_size", 1000) + self.chunk_overlap = kwargs.get("chunk_overlap", 200) # 元数据锁 self._metadata_lock = asyncio.Lock() @@ -73,11 +75,7 @@ class MilvusKB(KnowledgeBase): """初始化 Milvus 连接""" try: # 连接到 Milvus - connections.connect( - alias=self.connection_alias, - uri=self.milvus_uri, - token=self.milvus_token - ) + connections.connect(alias=self.connection_alias, uri=self.milvus_uri, token=self.milvus_token) # 创建数据库(如果不存在) try: @@ -106,10 +104,7 @@ class MilvusKB(KnowledgeBase): try: # 检查集合是否存在 if utility.has_collection(collection_name, using=self.connection_alias): - collection = Collection( - name=collection_name, - using=self.connection_alias - ) + collection = Collection(name=collection_name, using=self.connection_alias) # 检查嵌入模型是否匹配 description = collection.description @@ -137,27 +132,18 @@ class MilvusKB(KnowledgeBase): FieldSchema(name="chunk_id", dtype=DataType.VARCHAR, max_length=100), FieldSchema(name="file_id", dtype=DataType.VARCHAR, max_length=100), FieldSchema(name="chunk_index", dtype=DataType.INT64), - FieldSchema(name="embedding", dtype=DataType.FLOAT_VECTOR, dim=embedding_dim) + FieldSchema(name="embedding", dtype=DataType.FLOAT_VECTOR, dim=embedding_dim), ] schema = CollectionSchema( - fields=fields, - description=f"Knowledge base collection for {db_id} using {model_name}" + fields=fields, description=f"Knowledge base collection for {db_id} using {model_name}" ) # 创建集合 - collection = Collection( - name=collection_name, - schema=schema, - using=self.connection_alias - ) + collection = Collection(name=collection_name, schema=schema, using=self.connection_alias) # 创建索引 - index_params = { - "metric_type": "COSINE", - "index_type": "IVF_FLAT", - "params": {"nlist": 1024} - } + index_params = {"metric_type": "COSINE", "index_type": "IVF_FLAT", "params": {"nlist": 1024}} collection.create_index("embedding", index_params) logger.info(f"Created new Milvus collection: {collection_name}") @@ -218,18 +204,17 @@ class MilvusKB(KnowledgeBase): def _split_text_into_chunks(self, text: str, file_id: str, filename: str, params: dict) -> list[dict]: """将文本分割成块""" # 检查是否使用QA分割模式 - use_qa_split = params.get('use_qa_split', False) + use_qa_split = params.get("use_qa_split", False) if use_qa_split: # 使用QA分割模式 - qa_separator = params.get('qa_separator', '\n\n\n') + qa_separator = params.get("qa_separator", "\n\n\n") return split_text_into_qa_chunks(text, file_id, filename, qa_separator, params) else: # 使用传统分割模式 return split_text_into_chunks(text, file_id, filename, params) - async def add_content(self, db_id: str, items: list[str], - params: dict | None = {}) -> list[dict]: + async def add_content(self, db_id: str, items: list[str], params: dict | None = {}) -> list[dict]: """添加内容(文件/URL)""" if db_id not in self.databases_meta: raise ValueError(f"Database {db_id} not found") @@ -241,7 +226,7 @@ class MilvusKB(KnowledgeBase): embed_info = self.databases_meta[db_id].get("embed_info", {}) embedding_function = self._get_async_embedding_function(embed_info) - content_type = params.get('content_type', 'file') if params else 'file' + content_type = params.get("content_type", "file") if params else "file" processed_items_info = [] for item in items: @@ -280,7 +265,7 @@ class MilvusKB(KnowledgeBase): [chunk["chunk_id"] for chunk in chunks], [chunk["file_id"] for chunk in chunks], [chunk["chunk_index"] for chunk in chunks], - embeddings + embeddings, ] def _insert_and_flush(): @@ -294,7 +279,7 @@ class MilvusKB(KnowledgeBase): async with self._metadata_lock: self.files_meta[file_id]["status"] = "done" self._save_metadata() - file_record['status'] = "done" + file_record["status"] = "done" # 从处理队列中移除 self._remove_from_processing_queue(file_id) @@ -303,7 +288,7 @@ class MilvusKB(KnowledgeBase): async with self._metadata_lock: self.files_meta[file_id]["status"] = "failed" self._save_metadata() - file_record['status'] = "failed" + file_record["status"] = "failed" # 从处理队列中移除 self._remove_from_processing_queue(file_id) finally: @@ -334,7 +319,7 @@ class MilvusKB(KnowledgeBase): anns_field="embedding", param=search_params, limit=top_k, - output_fields=["content", "source", "chunk_id", "file_id", "chunk_index"] + output_fields=["content", "source", "chunk_id", "file_id", "chunk_index"], ) if not results or len(results) == 0 or len(results[0]) == 0: @@ -352,14 +337,12 @@ class MilvusKB(KnowledgeBase): "source": entity.get("source", "未知来源"), "chunk_id": entity.get("chunk_id"), "file_id": entity.get("file_id"), - "chunk_index": entity.get("chunk_index") + "chunk_index": entity.get("chunk_index"), } - retrieved_chunks.append({ - "content": entity.get("content", ""), - "metadata": metadata, - "score": similarity - }) + retrieved_chunks.append( + {"content": entity.get("content", ""), "metadata": metadata, "score": similarity} + ) logger.debug(f"Milvus query response: {len(retrieved_chunks)} chunks found (after similarity filtering)") return retrieved_chunks @@ -376,11 +359,7 @@ class MilvusKB(KnowledgeBase): # 先查询文件是否存在,避免不必要的删除操作 try: expr = f'file_id == "{file_id}"' - results = collection.query( - expr=expr, - output_fields=["id"], - limit=1 - ) + results = collection.query(expr=expr, output_fields=["id"], limit=1) if not results: logger.info(f"File {file_id} not found in Milvus, skipping delete operation") @@ -417,7 +396,7 @@ class MilvusKB(KnowledgeBase): results = collection.query( expr=expr, output_fields=["content", "chunk_id", "chunk_index"], - limit=10000 # 假设单个文件不会超过10000个chunks + limit=10000, # 假设单个文件不会超过10000个chunks ) # 构建chunks数据 @@ -426,7 +405,7 @@ class MilvusKB(KnowledgeBase): chunk_data = { "id": result.get("chunk_id", ""), "content": result.get("content", ""), - "chunk_order_index": result.get("chunk_index", 0) + "chunk_order_index": result.get("chunk_index", 0), } doc_chunks.append(chunk_data) @@ -457,7 +436,7 @@ class MilvusKB(KnowledgeBase): def __del__(self): """清理连接""" try: - if hasattr(self, 'connection_alias'): + if hasattr(self, "connection_alias"): connections.disconnect(self.connection_alias) except Exception: # noqa: S110 pass diff --git a/src/models/__init__.py b/src/models/__init__.py index a900c394..27b7ffbf 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -2,9 +2,10 @@ import os import traceback from src import config -from src.utils.logging_config import logger from src.models.chat_model import OpenAIBase from src.models.embedding import OllamaEmbedding, OtherEmbedding +from src.utils.logging_config import logger + def select_model(model_provider, model_name=None): """根据模型提供者选择模型""" @@ -16,12 +17,14 @@ def select_model(model_provider, model_name=None): if model_provider == "openai": from src.models.chat_model import OpenModel + return OpenModel(model_name) if model_provider == "custom": model_info = get_custom_model(model_name) from src.models.chat_model import CustomModel + return CustomModel(model_info) # 其他模型,默认使用OpenAIBase @@ -37,7 +40,7 @@ def select_model(model_provider, model_name=None): def select_embedding_model(model_id): - provider, model_name = model_id.split('/', 1) if model_id else ("", "") + provider, model_name = model_id.split("/", 1) if model_id else ("", "") support_embed_models = config.embed_model_names.keys() assert model_id in support_embed_models, f"Unsupported embed model: {model_id}, only support {support_embed_models}" logger.debug(f"Loading embedding model {model_id}") diff --git a/src/models/chat_model.py b/src/models/chat_model.py index 609173da..e70f68b0 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -1,8 +1,11 @@ import os + import requests -from openai import OpenAI -from src.utils import logger, get_docker_safe_url from langchain_openai import ChatOpenAI +from openai import OpenAI + +from src.utils import get_docker_safe_url, logger + class OpenAIBase: def __init__(self, api_key, base_url, model_name, **kwargs): @@ -14,7 +17,7 @@ class OpenAIBase: def predict(self, message, stream=False): if isinstance(message, str): - messages=[{"role": "user", "content": message}] + messages = [{"role": "user", "content": message}] else: messages = message @@ -31,8 +34,8 @@ class OpenAIBase: stream=True, ) for chunk in response: - if len(chunk.choices) > 0: - yield chunk.choices[0].delta + if len(chunk.choices) > 0: + yield chunk.choices[0].delta except Exception as e: err = f"Error streaming response: {e}, URL: {self.base_url}, API Key: {self.api_key[:5]}***, Model: {self.model_name}" @@ -49,11 +52,7 @@ class OpenAIBase: def get_models(self): try: - return self.client.models.list( - extra_query={ - "type": "text" - } - ) + return self.client.models.list(extra_query={"type": "text"}) except Exception as e: logger.error(f"Error getting models: {e}") return [] @@ -67,7 +66,6 @@ class OpenModel(OpenAIBase): super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) - class CustomModel(OpenAIBase): def __init__(self, model_info): model_name = model_info["name"] @@ -83,5 +81,6 @@ class GeneralResponse: self.content = content self.is_full = False + if __name__ == "__main__": pass diff --git a/src/models/embedding.py b/src/models/embedding.py index 9ae0626f..170f0c6c 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -1,16 +1,16 @@ -import os +import asyncio import json +import os +from abc import ABC, abstractmethod + import httpx import requests -import asyncio -from abc import abstractmethod, ABC from src import config -from src.utils import hashstr, logger, get_docker_safe_url +from src.utils import get_docker_safe_url, hashstr, logger class BaseEmbeddingModel(ABC): - def __init__(self, model=None, name=None, dimension=None, url=None, base_url=None, api_key=None): """ Args: @@ -60,22 +60,18 @@ class BaseEmbeddingModel(ABC): task_id = None if len(messages) > batch_size: task_id = hashstr(messages) - self.embed_state[task_id] = { - 'status': 'in-progress', - 'total': len(messages), - 'progress': 0 - } + self.embed_state[task_id] = {"status": "in-progress", "total": len(messages), "progress": 0} for i in range(0, len(messages), batch_size): - group_msg = messages[i:i+batch_size] + group_msg = messages[i : i + batch_size] logger.info(f"Encoding [{i}/{len(messages)}] messages (bsz={batch_size})") response = self.encode(group_msg) data.extend(response) if task_id: - self.embed_state[task_id]['progress'] = i + len(group_msg) + self.embed_state[task_id]["progress"] = i + len(group_msg) if task_id: - self.embed_state[task_id]['status'] = 'completed' + self.embed_state[task_id]["status"] = "completed" return data @@ -84,15 +80,11 @@ class BaseEmbeddingModel(ABC): task_id = None if len(messages) > batch_size: task_id = hashstr(messages) - self.embed_state[task_id] = { - 'status': 'in-progress', - 'total': len(messages), - 'progress': 0 - } + self.embed_state[task_id] = {"status": "in-progress", "total": len(messages), "progress": 0} tasks = [] for i in range(0, len(messages), batch_size): - group_msg = messages[i:i+batch_size] + group_msg = messages[i : i + batch_size] tasks.append(self.aencode(group_msg)) results = await asyncio.gather(*tasks) @@ -100,11 +92,12 @@ class BaseEmbeddingModel(ABC): data.extend(res) if task_id: - self.embed_state[task_id]['progress'] = len(messages) - self.embed_state[task_id]['status'] = 'completed' + self.embed_state[task_id]["progress"] = len(messages) + self.embed_state[task_id]["status"] = "completed" return data + class OllamaEmbedding(BaseEmbeddingModel): """ Ollama Embedding Model @@ -149,13 +142,9 @@ class OllamaEmbedding(BaseEmbeddingModel): class OtherEmbedding(BaseEmbeddingModel): - def __init__(self, **kwargs) -> None: super().__init__(**kwargs) - self.headers = { - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/json" - } + self.headers = {"Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json"} def build_payload(self, message: list[str] | str) -> dict: return {"model": self.model, "input": message} diff --git a/src/models/rerank_model.py b/src/models/rerank_model.py index f51944d7..574b0e62 100644 --- a/src/models/rerank_model.py +++ b/src/models/rerank_model.py @@ -1,25 +1,25 @@ -import os import json -import requests +import os + import numpy as np +import requests from src import config -from src.utils import logger, get_docker_safe_url +from src.utils import get_docker_safe_url, logger + def sigmoid(x): return 1 / (1 + np.exp(-x)) + class OnlineReranker: def __init__(self, model_name, api_key, base_url, **kwargs): self.url = get_docker_safe_url(base_url) self.model = model_name self.api_key = api_key - self.headers = { - "Authorization": f"Bearer {api_key}", - "Content-Type": "application/json" - } + self.headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"} - def compute_score(self, sentence_pairs, batch_size = 256, max_length = 512, normalize = False): + def compute_score(self, sentence_pairs, batch_size=256, max_length=512, normalize=False): # TODO 还没实现 batch_size query, sentences = sentence_pairs[0], sentence_pairs[1] payload = self.build_payload(query, sentences, max_length) @@ -35,7 +35,7 @@ class OnlineReranker: return all_scores - def build_payload(self, query, sentences, max_length = 512): + def build_payload(self, query, sentences, max_length=512): return { "model": self.model, "query": query, @@ -43,6 +43,7 @@ class OnlineReranker: "max_chunks_per_doc": max_length, } + def get_reranker(model_id, **kwargs): support_rerankers = config.reranker_names.keys() assert model_id in support_rerankers, f"Unsupported Reranker: {model_id}, only support {support_rerankers}" @@ -51,9 +52,4 @@ def get_reranker(model_id, **kwargs): base_url = model_info["base_url"] api_key = os.getenv(model_info["api_key"], model_info["api_key"]) assert api_key, f"{model_info['name']} api_key is required" - return OnlineReranker( - model_name=model_info["name"], - api_key=api_key, - base_url=base_url, - **kwargs - ) + return OnlineReranker(model_name=model_info["name"], api_key=api_key, base_url=base_url, **kwargs) diff --git a/src/plugins/_ocr.py b/src/plugins/_ocr.py index 535a4789..1b387693 100644 --- a/src/plugins/_ocr.py +++ b/src/plugins/_ocr.py @@ -1,27 +1,22 @@ import os -import uuid import time -from pathlib import Path +import uuid from argparse import ArgumentParser from collections import defaultdict +from pathlib import Path import fitz # fitz就是pip install PyMuPDF import numpy as np # Added import for numpy from PIL import Image -from tqdm import tqdm from rapidocr_onnxruntime import RapidOCR +from tqdm import tqdm -from src.utils import logger, is_text_pdf - +from src.utils import is_text_pdf, logger GOLBAL_STATE = {} # OCR服务监控统计 -OCR_STATS = { - "requests": defaultdict(int), - "failures": defaultdict(int), - "service_status": defaultdict(str) -} +OCR_STATS = {"requests": defaultdict(int), "failures": defaultdict(int), "service_status": defaultdict(str)} def log_ocr_request(service_name: str, file_path: str, success: bool, processing_time: float, error_msg: str = None): @@ -50,7 +45,7 @@ def get_ocr_stats(): "success_count": success_count, "failure_count": OCR_STATS["failures"][service], "success_rate": f"{success_rate:.2%}", - "status": OCR_STATS["service_status"][service] + "status": OCR_STATS["service_status"][service], } return stats @@ -58,6 +53,7 @@ def get_ocr_stats(): class OCRServiceException(Exception): """OCR服务异常""" + def __init__(self, message, service_name=None, status_code=None): super().__init__(message) self.service_name = service_name @@ -69,8 +65,10 @@ class OCRPlugin: def __init__(self, **kwargs): self.ocr = None - self.det_box_thresh = kwargs.get('det_box_thresh', 0.3) - self.model_dir_root = os.getenv("MODEL_DIR") if not os.getenv("RUNNING_IN_DOCKER") else os.getenv("MODEL_DIR_IN_DOCKER") + self.det_box_thresh = kwargs.get("det_box_thresh", 0.3) + self.model_dir_root = ( + os.getenv("MODEL_DIR") if not os.getenv("RUNNING_IN_DOCKER") else os.getenv("MODEL_DIR_IN_DOCKER") + ) def _check_rapid_ocr_availability(self): """检查RapidOCR模型是否可用""" @@ -81,16 +79,14 @@ class OCRPlugin: if not os.path.exists(model_dir): raise OCRServiceException( - f"模型目录不存在: {model_dir}。请下载 SWHL/RapidOCR 模型", - "rapid_ocr", - "model_not_found" + f"模型目录不存在: {model_dir}。请下载 SWHL/RapidOCR 模型", "rapid_ocr", "model_not_found" ) if not os.path.exists(det_model_dir) or not os.path.exists(rec_model_dir): raise OCRServiceException( f"模型文件缺失。请确认模型文件完整: {det_model_dir}, {rec_model_dir}", "rapid_ocr", - "model_incomplete" + "model_incomplete", ) return True @@ -99,11 +95,7 @@ class OCRPlugin: if isinstance(e, OCRServiceException): raise else: - raise OCRServiceException( - f"RapidOCR模型检查失败: {str(e)}", - "rapid_ocr", - "check_failed" - ) + raise OCRServiceException(f"RapidOCR模型检查失败: {str(e)}", "rapid_ocr", "check_failed") def load_model(self): """加载 OCR 模型""" @@ -120,11 +112,7 @@ class OCRPlugin: self.ocr = RapidOCR(det_box_thresh=0.3, det_model_path=det_model_dir, rec_model_path=rec_model_dir) logger.info(f"OCR Plugin for det_box_thresh = {self.det_box_thresh} loaded.") except Exception as e: - raise OCRServiceException( - f"RapidOCR模型加载失败: {str(e)}", - "rapid_ocr", - "load_failed" - ) + raise OCRServiceException(f"RapidOCR模型加载失败: {str(e)}", "rapid_ocr", "load_failed") def process_image(self, image, params=None): """ @@ -165,7 +153,7 @@ class OCRPlugin: # 提取文本 if result: - text = '\n'.join([line[1] for line in result]) + text = "\n".join([line[1] for line in result]) log_ocr_request("rapid_ocr", image_path, True, processing_time) return text else: @@ -189,11 +177,11 @@ class OCRPlugin: str: 临时文件路径 """ # 为临时文件创建目录(如果不存在) - tmp_dir = os.path.join(os.getcwd(), 'tmp') + tmp_dir = os.path.join(os.getcwd(), "tmp") os.makedirs(tmp_dir, exist_ok=True) # 生成临时文件路径 - temp_filename = f'ocr_temp_{uuid.uuid4().hex[:8]}.png' + temp_filename = f"ocr_temp_{uuid.uuid4().hex[:8]}.png" image_path = os.path.join(tmp_dir, temp_filename) # 根据图像类型保存文件 @@ -224,7 +212,7 @@ class OCRPlugin: pdfDoc = fitz.open(pdf_path) totalPage = pdfDoc.page_count - for pg in tqdm(range(totalPage), desc='to images', ncols=100): + for pg in tqdm(range(totalPage), desc="to images", ncols=100): page = pdfDoc[pg] rotate, zoom_x, zoom_y = 0, 2, 2 mat = fitz.Matrix(zoom_x, zoom_y).prerotate(rotate) @@ -234,12 +222,12 @@ class OCRPlugin: # 处理每个图像并合并文本 all_text = [] - for img_path in tqdm(images, desc='to txt', ncols=100): + for img_path in tqdm(images, desc="to txt", ncols=100): text = self.process_image(img_path) all_text.append(text) logger.debug(f"PDF OCR result: {all_text[:50]}(...) total {len(all_text)} pages.") - return '\n\n'.join(all_text) + return "\n\n".join(all_text) except Exception as e: logger.error(f"PDF processing error: {str(e)}") @@ -253,6 +241,7 @@ class OCRPlugin: :return: 提取的文本 """ import requests + from .mineru import parse_doc mineru_ocr_uri = os.getenv("MINERU_OCR_URI", "http://localhost:30000") @@ -269,28 +258,20 @@ class OCRPlugin: error_detail = health_check_response.text raise OCRServiceException( - f"MinerU OCR服务健康检查失败: {error_detail}", - "mineru_ocr", - "health_check_failed" + f"MinerU OCR服务健康检查失败: {error_detail}", "mineru_ocr", "health_check_failed" ) except Exception as e: if isinstance(e, OCRServiceException): raise - raise OCRServiceException( - f"MinerU OCR服务检查失败: {str(e)}", - "mineru_ocr", - "service_error" - ) + raise OCRServiceException(f"MinerU OCR服务检查失败: {str(e)}", "mineru_ocr", "service_error") try: start_time = time.time() file_path_list = [file_path] output_dir = os.path.join(os.getcwd(), "tmp", "mineru_ocr") - text = parse_doc(file_path_list, output_dir, - backend="vlm-sglang-client", - server_url=mineru_ocr_uri)[0] + text = parse_doc(file_path_list, output_dir, backend="vlm-sglang-client", server_url=mineru_ocr_uri)[0] processing_time = time.time() - start_time log_ocr_request("mineru_ocr", file_path, True, processing_time) @@ -303,11 +284,7 @@ class OCRPlugin: error_msg = f"MinerU OCR处理失败: {str(e)}" log_ocr_request("mineru_ocr", file_path, False, processing_time, error_msg) - raise OCRServiceException( - error_msg, - "mineru_ocr", - "processing_failed" - ) + raise OCRServiceException(error_msg, "mineru_ocr", "processing_failed") def process_file_paddlex(self, pdf_path, params=None): """ @@ -331,18 +308,12 @@ class OCRPlugin: error_detail = health_check_response.text raise OCRServiceException( - f"PaddleX OCR服务健康检查失败: {error_detail}", - "paddlex_ocr", - "health_check_failed" + f"PaddleX OCR服务健康检查失败: {error_detail}", "paddlex_ocr", "health_check_failed" ) except Exception as e: if isinstance(e, OCRServiceException): raise - raise OCRServiceException( - f"PaddleX OCR服务检查失败: {str(e)}", - "paddlex_ocr", - "service_error" - ) + raise OCRServiceException(f"PaddleX OCR服务检查失败: {str(e)}", "paddlex_ocr", "service_error") try: start_time = time.time() @@ -353,11 +324,7 @@ class OCRPlugin: error_msg = f"PaddleX OCR处理失败: {result['error']}" log_ocr_request("paddlex_ocr", pdf_path, False, processing_time, error_msg) - raise OCRServiceException( - error_msg, - "paddlex_ocr", - "processing_failed" - ) + raise OCRServiceException(error_msg, "paddlex_ocr", "processing_failed") log_ocr_request("paddlex_ocr", pdf_path, True, processing_time) return result["full_text"] @@ -365,19 +332,17 @@ class OCRPlugin: except Exception as e: if isinstance(e, OCRServiceException): raise - processing_time = time.time() - start_time if 'start_time' in locals() else 0 + processing_time = time.time() - start_time if "start_time" in locals() else 0 error_msg = f"PaddleX OCR处理失败: {str(e)}" log_ocr_request("paddlex_ocr", pdf_path, False, processing_time, error_msg) - raise OCRServiceException( - error_msg, - "paddlex_ocr", - "processing_failed" - ) + raise OCRServiceException(error_msg, "paddlex_ocr", "processing_failed") + def get_state(task_id): return GOLBAL_STATE.get(task_id, {}) + def plainreader(file_path): """读取普通文本文件并返回text文本""" assert os.path.exists(file_path), "File not found" @@ -389,8 +354,8 @@ def plainreader(file_path): if __name__ == "__main__": parser = ArgumentParser() - parser.add_argument('--pdf-path', type=str, required=True, help='Path to the PDF file') - parser.add_argument('--return-text', action='store_true', help='Return the extracted text') + parser.add_argument("--pdf-path", type=str, required=True, help="Path to the PDF file") + parser.add_argument("--return-text", action="store_true", help="Return the extracted text") args = parser.parse_args() ocr = OCRPlugin() diff --git a/src/plugins/mineru.py b/src/plugins/mineru.py index 67d3bbc7..6acd3cfa 100644 --- a/src/plugins/mineru.py +++ b/src/plugins/mineru.py @@ -3,17 +3,17 @@ import copy import json import os from pathlib import Path -from tqdm import tqdm +from mineru.backend.pipeline.model_json_to_middle_json import result_to_middle_json as pipeline_result_to_middle_json +from mineru.backend.pipeline.pipeline_analyze import doc_analyze as pipeline_doc_analyze +from mineru.backend.pipeline.pipeline_middle_json_mkcontent import union_make as pipeline_union_make +from mineru.backend.vlm.vlm_analyze import doc_analyze as vlm_doc_analyze +from mineru.backend.vlm.vlm_middle_json_mkcontent import union_make as vlm_union_make from mineru.cli.common import convert_pdf_bytes_to_bytes_by_pypdfium2, prepare_env, read_fn from mineru.data.data_reader_writer import FileBasedDataWriter from mineru.utils.draw_bbox import draw_layout_bbox, draw_span_bbox from mineru.utils.enum_class import MakeMode -from mineru.backend.vlm.vlm_analyze import doc_analyze as vlm_doc_analyze -from mineru.backend.pipeline.pipeline_analyze import doc_analyze as pipeline_doc_analyze -from mineru.backend.pipeline.pipeline_middle_json_mkcontent import union_make as pipeline_union_make -from mineru.backend.pipeline.model_json_to_middle_json import result_to_middle_json as pipeline_result_to_middle_json -from mineru.backend.vlm.vlm_middle_json_mkcontent import union_make as vlm_union_make +from tqdm import tqdm from src.utils.logging_config import logger @@ -39,17 +39,20 @@ def do_parse( start_page_id=0, # Start page ID for parsing, default is 0 end_page_id=None, # End page ID for parsing, default is None (parse all pages until the end of the document) ) -> list[str]: - if backend == "pipeline": for idx, pdf_bytes in enumerate(pdf_bytes_list): new_pdf_bytes = convert_pdf_bytes_to_bytes_by_pypdfium2(pdf_bytes, start_page_id, end_page_id) pdf_bytes_list[idx] = new_pdf_bytes - - result = pipeline_doc_analyze(pdf_bytes_list, p_lang_list, parse_method=parse_method, formula_enable=p_formula_enable,table_enable=p_table_enable) + result = pipeline_doc_analyze( + pdf_bytes_list, + p_lang_list, + parse_method=parse_method, + formula_enable=p_formula_enable, + table_enable=p_table_enable, + ) infer_results, all_image_lists, all_pdf_docs, lang_list, ocr_enabled_list = result - md_results = [] for idx, model_list in enumerate(infer_results): model_json = copy.deepcopy(model_list) @@ -61,7 +64,9 @@ def do_parse( pdf_doc = all_pdf_docs[idx] _lang = lang_list[idx] _ocr_enable = ocr_enabled_list[idx] - middle_json = pipeline_result_to_middle_json(model_list, images_list, pdf_doc, image_writer, _lang, _ocr_enable, p_formula_enable) + middle_json = pipeline_result_to_middle_json( + model_list, images_list, pdf_doc, image_writer, _lang, _ocr_enable, p_formula_enable + ) pdf_info = middle_json["pdf_info"] @@ -124,7 +129,9 @@ def do_parse( pdf_bytes = convert_pdf_bytes_to_bytes_by_pypdfium2(pdf_bytes, start_page_id, end_page_id) local_image_dir, local_md_dir = prepare_env(output_dir, pdf_file_name, parse_method) image_writer, md_writer = FileBasedDataWriter(local_image_dir), FileBasedDataWriter(local_md_dir) - middle_json, infer_result = vlm_doc_analyze(pdf_bytes, image_writer=image_writer, backend=backend, server_url=server_url) + middle_json, infer_result = vlm_doc_analyze( + pdf_bytes, image_writer=image_writer, backend=backend, server_url=server_url + ) pdf_info = middle_json["pdf_info"] @@ -176,35 +183,35 @@ def do_parse( def parse_doc( - path_list: list[Path], - output_dir, - lang="ch", - backend="pipeline", - method="auto", - server_url=None, - start_page_id=0, # Start page ID for parsing, default is 0 - end_page_id=None # End page ID for parsing, default is None (parse all pages until the end of the document) + path_list: list[Path], + output_dir, + lang="ch", + backend="pipeline", + method="auto", + server_url=None, + start_page_id=0, # Start page ID for parsing, default is 0 + end_page_id=None, # End page ID for parsing, default is None (parse all pages until the end of the document) ) -> list[str]: """ - Parameter description: - path_list: List of document paths to be parsed, can be PDF or image files. - output_dir: Output directory for storing parsing results. - lang: Language option, default is 'ch', optional values include['ch', 'ch_server', 'ch_lite', 'en', 'korean', 'japan', 'chinese_cht', 'ta', 'te', 'ka']。 - Input the languages in the pdf (if known) to improve OCR accuracy. Optional. - Adapted only for the case where the backend is set to "pipeline" - backend: the backend for parsing pdf: - pipeline: More general. - vlm-transformers: More general. - vlm-sglang-engine: Faster(engine). - vlm-sglang-client: Faster(client). - without method specified, pipeline will be used by default. - method: the method for parsing pdf: - auto: Automatically determine the method based on the file type. - txt: Use text extraction method. - ocr: Use OCR method for image-based PDFs. - Without method specified, 'auto' will be used by default. - Adapted only for the case where the backend is set to "pipeline". - server_url: When the backend is `sglang-client`, you need to specify the server_url, for example:`http://127.0.0.1:30000` + Parameter description: + path_list: List of document paths to be parsed, can be PDF or image files. + output_dir: Output directory for storing parsing results. + lang: Language option, default is 'ch', optional values include['ch', 'ch_server', 'ch_lite', 'en', 'korean', 'japan', 'chinese_cht', 'ta', 'te', 'ka']。 + Input the languages in the pdf (if known) to improve OCR accuracy. Optional. + Adapted only for the case where the backend is set to "pipeline" + backend: the backend for parsing pdf: + pipeline: More general. + vlm-transformers: More general. + vlm-sglang-engine: Faster(engine). + vlm-sglang-client: Faster(client). + without method specified, pipeline will be used by default. + method: the method for parsing pdf: + auto: Automatically determine the method based on the file type. + txt: Use text extraction method. + ocr: Use OCR method for image-based PDFs. + Without method specified, 'auto' will be used by default. + Adapted only for the case where the backend is set to "pipeline". + server_url: When the backend is `sglang-client`, you need to specify the server_url, for example:`http://127.0.0.1:30000` """ file_name_list = [] pdf_bytes_list = [] @@ -225,21 +232,22 @@ def parse_doc( parse_method=method, server_url=server_url, start_page_id=start_page_id, - end_page_id=end_page_id + end_page_id=end_page_id, ) return result if result else [""] if __name__ == "__main__": - pdf_files_dir = "/home/zwj/workspace/projects/Yuxi-Know/test/struct_pdf" output_dir = "/home/zwj/workspace/projects/Yuxi-Know/test/struct_pdf_output" pdf_suffixes = [".pdf"] image_suffixes = [".png", ".jpeg", ".jpg"] doc_path_list = [] - for doc_path in Path(pdf_files_dir).glob('*'): + for doc_path in Path(pdf_files_dir).glob("*"): if doc_path.suffix in pdf_suffixes + image_suffixes: doc_path_list.append(doc_path) - parse_doc(doc_path_list, output_dir, backend="vlm-sglang-client", server_url="http://172.19.13.5:30000") # faster(client). + parse_doc( + doc_path_list, output_dir, backend="vlm-sglang-client", server_url="http://172.19.13.5:30000" + ) # faster(client). diff --git a/src/plugins/paddlex.py b/src/plugins/paddlex.py index 3e26a854..5fff01c0 100644 --- a/src/plugins/paddlex.py +++ b/src/plugins/paddlex.py @@ -1,30 +1,29 @@ - -import requests -import json import base64 +import json import os import time -from typing import Optional, Any from pathlib import Path +from typing import Any, Optional + +import requests if __name__ == "__main__": - from loguru import logger import typer + from loguru import logger else: from src.utils import logger - class PaddleXLayoutParser: """PaddleX 版面解析服务客户端""" def __init__(self, base_url: str = "http://localhost:8080"): - self.base_url = base_url.rstrip('/') + self.base_url = base_url.rstrip("/") self.endpoint = f"{self.base_url}/layout-parsing" def encode_file_to_base64(self, file_path: str) -> str: - with open(file_path, 'rb') as file: - encoded = base64.b64encode(file.read()).decode('utf-8') + with open(file_path, "rb") as file: + encoded = base64.b64encode(file.read()).decode("utf-8") return encoded def _process_file_input(self, file_input: str) -> str: @@ -43,7 +42,7 @@ class PaddleXLayoutParser: raise # 检查是否为URL - elif file_input.startswith(('http://', 'https://')): + elif file_input.startswith(("http://", "https://")): logger.info(f"🌐 检测到URL: {file_input}") return file_input @@ -52,21 +51,23 @@ class PaddleXLayoutParser: logger.info(f"📝 假设为Base64编码内容,长度: {len(file_input)} 字符") return file_input - def layout_parsing(self, - file_input: str, - file_type: int | None = None, - use_textline_orientation: bool | None = None, - use_seal_recognition: bool | None = None, - use_table_recognition: bool | None = None, - use_formula_recognition: bool | None = None, - use_chart_recognition: bool | None = None, - use_region_detection: bool | None = None, - layout_threshold: float | None = None, - layout_nms: bool | None = None, - use_doc_orientation_classify: bool = True, - use_doc_unwarping: bool | None = False, - use_wired_table_cells_trans_to_html: bool = True, # 启用则直接基于有线表单元格检测结果的几何关系构建HTML。 - **kwargs) -> dict[str, Any]: + def layout_parsing( + self, + file_input: str, + file_type: int | None = None, + use_textline_orientation: bool | None = None, + use_seal_recognition: bool | None = None, + use_table_recognition: bool | None = None, + use_formula_recognition: bool | None = None, + use_chart_recognition: bool | None = None, + use_region_detection: bool | None = None, + layout_threshold: float | None = None, + layout_nms: bool | None = None, + use_doc_orientation_classify: bool = True, + use_doc_unwarping: bool | None = False, + use_wired_table_cells_trans_to_html: bool = True, # 启用则直接基于有线表单元格检测结果的几何关系构建HTML。 + **kwargs, + ) -> dict[str, Any]: """ 调用版面解析API:https://paddlepaddle.github.io/PaddleX/latest/pipeline_usage/tutorials/ocr_pipelines/PP-StructureV3.html#22-python """ @@ -102,10 +103,7 @@ class PaddleXLayoutParser: try: response = requests.post( - self.endpoint, - json=payload, - headers={"Content-Type": "application/json"}, - timeout=300 + self.endpoint, json=payload, headers={"Content-Type": "application/json"}, timeout=300 ) if response.status_code == 200: @@ -132,7 +130,6 @@ class PaddleXLayoutParser: return {"error": str(e)} - def _parse_recognition_result(api_result: dict[str, Any], file_path: str) -> dict[str, Any]: # 基本信息 parsed_result = { @@ -143,7 +140,7 @@ def _parse_recognition_result(api_result: dict[str, Any], file_path: str) -> dic "total_pages": 0, "pages": [], "full_text": "", - "summary": {} + "summary": {}, } result_data = api_result.get("result", {}) @@ -155,7 +152,7 @@ def _parse_recognition_result(api_result: dict[str, Any], file_path: str) -> dic parsed_result["document_info"] = { "type": data_info.get("type", "unknown"), "total_pages": data_info.get("numPages", len(layout_results)), - "page_dimensions": data_info.get("pages", []) + "page_dimensions": data_info.get("pages", []), } # 统计信息 @@ -168,11 +165,7 @@ def _parse_recognition_result(api_result: dict[str, Any], file_path: str) -> dic # 解析每页结果 for page_index, page_result in enumerate(layout_results): - page_info = { - "page_number": page_index + 1, - "content": {}, - "statistics": {} - } + page_info = {"page_number": page_index + 1, "content": {}, "statistics": {}} # Markdown内容 if "markdown" in page_result: @@ -240,37 +233,28 @@ def _parse_recognition_result(api_result: dict[str, Any], file_path: str) -> dic "total_charts": total_charts, "total_seals": total_seals, "total_characters": len(parsed_result["full_text"]), - "average_elements_per_page": round(total_elements / max(1, len(layout_results)), 2) + "average_elements_per_page": round(total_elements / max(1, len(layout_results)), 2), } return parsed_result def analyze_document(file_path: str, base_url: str = "http://localhost:8080") -> dict[str, Any]: - # 检查文件是否存在 if not os.path.exists(file_path): - return { - "success": False, - "error": f"文件不存在: {file_path}", - "file_path": file_path - } + return {"success": False, "error": f"文件不存在: {file_path}", "file_path": file_path} # 初始化客户端 client = PaddleXLayoutParser(base_url=base_url) # 判断文件类型 file_ext = os.path.splitext(file_path)[1].lower() - if file_ext == '.pdf': + if file_ext == ".pdf": file_type = 0 - elif file_ext in ['.jpg', '.jpeg', '.png', '.bmp', '.tiff', '.tif']: + elif file_ext in [".jpg", ".jpeg", ".png", ".bmp", ".tiff", ".tif"]: file_type = 1 else: - return { - "success": False, - "error": f"不支持的文件类型: {file_ext}", - "file_path": file_path - } + return {"success": False, "error": f"不支持的文件类型: {file_ext}", "file_path": file_path} logger.info(f"📄 开始分析文档: {os.path.basename(file_path)}") logger.info(f"📏 文件大小: {os.path.getsize(file_path) / 1024 / 1024:.2f} MB") @@ -286,7 +270,7 @@ def analyze_document(file_path: str, base_url: str = "http://localhost:8080") -> "success": False, "error": result.get("errorMsg", "API调用失败"), "file_path": file_path, - "raw_result": result + "raw_result": result, } # 解析结果 @@ -294,11 +278,7 @@ def analyze_document(file_path: str, base_url: str = "http://localhost:8080") -> return analysis_result except Exception as e: - return { - "success": False, - "error": f"处理异常: {str(e)}", - "file_path": file_path - } + return {"success": False, "error": f"处理异常: {str(e)}", "file_path": file_path} def check_paddlex_health(base_url: str = "http://localhost:8080") -> bool: @@ -318,7 +298,7 @@ def analyze_folder(input_dir: str, output_dir: str, base_url: str = "http://loca output_path.mkdir(parents=True, exist_ok=True) # 获取所有支持的文件 - supported_extensions = {'.pdf', '.jpg', '.jpeg', '.png', '.bmp', '.tiff', '.tif'} + supported_extensions = {".pdf", ".jpg", ".jpeg", ".png", ".bmp", ".tiff", ".tif"} files = [] for root, dirs, filenames in os.walk(input_dir): for filename in filenames: @@ -343,12 +323,16 @@ def analyze_folder(input_dir: str, output_dir: str, base_url: str = "http://loca if result.get("success"): # 保持目录结构 relative_path = file_path.relative_to(input_path) - output_file = output_path / relative_path.with_suffix('.txt') + output_file = output_path / relative_path.with_suffix(".txt") output_file.parent.mkdir(parents=True, exist_ok=True) # 写入文本内容 - text_content = result.get("full_text", "未提取到内容") if result.get("success") else f"分析失败:{result.get('error', '未知错误')}" - with open(output_file, 'w', encoding='utf-8') as f: + text_content = ( + result.get("full_text", "未提取到内容") + if result.get("success") + else f"分析失败:{result.get('error', '未知错误')}" + ) + with open(output_file, "w", encoding="utf-8") as f: f.write(text_content) success_count += 1 diff --git a/src/utils/__init__.py b/src/utils/__init__.py index b1c0da2d..40ca0df7 100644 --- a/src/utils/__init__.py +++ b/src/utils/__init__.py @@ -1,10 +1,13 @@ -import time import hashlib import os +import time + from src.utils.logging_config import logger + def is_text_pdf(pdf_path): import fitz + doc = fitz.open(pdf_path) total_pages = len(doc) if total_pages == 0: @@ -22,6 +25,7 @@ def is_text_pdf(pdf_path): # 如果超过50%的页面有文本内容,则认为是文本PDF return text_ratio > 0.5 + def hashstr(input_string, length=None, with_salt=False): """生成字符串的哈希值 Args: @@ -31,14 +35,14 @@ def hashstr(input_string, length=None, with_salt=False): """ try: # 尝试直接编码 - encoded_string = str(input_string).encode('utf-8') + encoded_string = str(input_string).encode("utf-8") except UnicodeEncodeError: # 如果编码失败,替换无效字符 - encoded_string = str(input_string).encode('utf-8', errors='replace') + encoded_string = str(input_string).encode("utf-8", errors="replace") if with_salt: salt = str(time.time()) - encoded_string = (encoded_string.decode('utf-8') + salt).encode('utf-8') + encoded_string = (encoded_string.decode("utf-8") + salt).encode("utf-8") hash = hashlib.md5(encoded_string).hexdigest() if length: diff --git a/src/utils/logging_config.py b/src/utils/logging_config.py index d1563f29..c3e089c3 100644 --- a/src/utils/logging_config.py +++ b/src/utils/logging_config.py @@ -2,12 +2,12 @@ import os from datetime import datetime import pytz - from loguru import logger as loguru_logger -SAVE_DIR = os.getenv('SAVE_DIR', 'saves') -DATETIME = datetime.now(pytz.timezone('Asia/Shanghai')).strftime('%Y-%m-%d') -LOG_FILE = f'{SAVE_DIR}/logs/yuxi-{DATETIME}.log' +SAVE_DIR = os.getenv("SAVE_DIR", "saves") +DATETIME = datetime.now(pytz.timezone("Asia/Shanghai")).strftime("%Y-%m-%d") +LOG_FILE = f"{SAVE_DIR}/logs/yuxi-{DATETIME}.log" + def setup_logger(name, level="DEBUG", console=True): """使用 loguru 设置日志记录器""" @@ -24,7 +24,7 @@ def setup_logger(name, level="DEBUG", console=True): encoding="utf-8", rotation="10 MB", # 文件大小达到 10MB 时轮转 retention="30 days", # 保留30天的日志 - compression="zip" # 压缩旧日志文件 + compression="zip", # 压缩旧日志文件 ) # 添加控制台日志(有颜色) @@ -33,16 +33,16 @@ def setup_logger(name, level="DEBUG", console=True): lambda msg: print(msg, end=""), level=level, format="{time:MM-DD HH:mm:ss} {level} {name}:{line}: {message}", - colorize=True + colorize=True, ) return loguru_logger # 设置根日志记录器 -logger = setup_logger('Yuxi') +logger = setup_logger("Yuxi") -__all__ = ['logger'] +__all__ = ["logger"] # If you want to disable logging from external libraries # logging.getLogger('some_external_library').setLevel(logging.CRITICAL) diff --git a/src/utils/prompts.py b/src/utils/prompts.py index 540d2193..be740692 100644 --- a/src/utils/prompts.py +++ b/src/utils/prompts.py @@ -1,7 +1,8 @@ from datetime import datetime + def get_system_prompt(): - return (f"当前时间:{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n") + return f"当前时间:{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n" knowbase_qa_template = """ @@ -65,14 +66,4 @@ keywords_prompt_template = """ <文本>{text} """ -HYDE_PROMPT_TEMPLATE = ( - "Please write a passage to answer the question\n" - "Try to include as many key details as possible.\n" - "\n" - "\n" - "{context_str}\n" - "\n" - "{query}\n" - "\n" - 'Passage:\n' -) +HYDE_PROMPT_TEMPLATE = "Please write a passage to answer the question\nTry to include as many key details as possible.\n\n\n{context_str}\n\n{query}\n\nPassage:\n" diff --git a/src/utils/web_search.py b/src/utils/web_search.py index d8b0e3c4..4fd10671 100644 --- a/src/utils/web_search.py +++ b/src/utils/web_search.py @@ -1,7 +1,10 @@ import os + from tavily import TavilyClient + from src.utils.logging_config import logger + class WebSearcher: def __init__(self): api_key = os.getenv("TAVILY_API_KEY") @@ -22,21 +25,19 @@ class WebSearcher: 搜索结果列表 """ try: - search_results = self.client.search( - query=query, - search_depth="basic", - max_results=max_results - ) + search_results = self.client.search(query=query, search_depth="basic", max_results=max_results) # 提取需要的信息 formatted_results = [] - for result in search_results['results'][:max_results]: - formatted_results.append({ - 'title': result.get('title', ''), - 'content': result.get('content', ''), - 'url': result.get('url', ''), - 'score': result.get('score', 0) - }) + for result in search_results["results"][:max_results]: + formatted_results.append( + { + "title": result.get("title", ""), + "content": result.get("content", ""), + "url": result.get("url", ""), + "score": result.get("score", 0), + } + ) return formatted_results diff --git a/test/test_concurrency.py b/test/test_concurrency.py index 0061d3bb..5a78b9a6 100644 --- a/test/test_concurrency.py +++ b/test/test_concurrency.py @@ -1,15 +1,13 @@ import asyncio -import aiohttp import time +import aiohttp + async def make_request(session: aiohttp.ClientSession, request_id: int) -> dict: """发送单个请求到API""" url = "http://localhost:5000/chat/call" - payload = { - "query": "写一个冒泡排序", - "meta": {} - } + payload = {"query": "写一个冒泡排序", "meta": {}} start_time = time.time() print(f"请求 {request_id} 开始时间: {time.strftime('%H:%M:%S', time.localtime(start_time))}") @@ -19,19 +17,23 @@ async def make_request(session: aiohttp.ClientSession, request_id: int) -> dict: print(f"请求 {request_id} 结果: {result}") end_time = time.time() duration = end_time - start_time - print(f"请求 {request_id} 完成时间: {time.strftime('%H:%M:%S', time.localtime(end_time))} (耗时: {duration:.2f}秒)") + print( + f"请求 {request_id} 完成时间: {time.strftime('%H:%M:%S', time.localtime(end_time))} (耗时: {duration:.2f}秒)" + ) return { "request_id": request_id, "status": response.status, "time": duration, "start_time": start_time, "end_time": end_time, - "success": True + "success": True, } except Exception as e: end_time = time.time() duration = end_time - start_time - print(f"请求 {request_id} 失败时间: {time.strftime('%H:%M:%S', time.localtime(end_time))} (耗时: {duration:.2f}秒)") + print( + f"请求 {request_id} 失败时间: {time.strftime('%H:%M:%S', time.localtime(end_time))} (耗时: {duration:.2f}秒)" + ) return { "request_id": request_id, "status": None, @@ -39,15 +41,17 @@ async def make_request(session: aiohttp.ClientSession, request_id: int) -> dict: "start_time": start_time, "end_time": end_time, "success": False, - "error": str(e) + "error": str(e), } + async def run_concurrent_test(num_requests: int = 10) -> list[dict]: """运行并发测试""" async with aiohttp.ClientSession() as session: tasks = [make_request(session, i) for i in range(num_requests)] return await asyncio.gather(*tasks) + def analyze_results(results: list[dict]) -> None: """分析并打印测试结果""" total_requests = len(results) @@ -97,9 +101,9 @@ def analyze_results(results: list[dict]) -> None: active_requests = [t for t in active_requests if t > start_time] active_requests.append(end_time) - print(f"{result['request_id']:^7} {time.strftime('%H:%M:%S', time.localtime(start_time))} " - f"{time.strftime('%H:%M:%S', time.localtime(end_time))} " - f"{result['time']:^8.2f} {len(active_requests):^6}") + print( + f"{result['request_id']:^7} {time.strftime('%H:%M:%S', time.localtime(start_time))} {time.strftime('%H:%M:%S', time.localtime(end_time))} {result['time']:^8.2f} {len(active_requests):^6}" + ) # 计算最大并发数 max_concurrent = 0 @@ -116,6 +120,7 @@ def analyze_results(results: list[dict]) -> None: print(f"\n最大并发请求数: {max_concurrent}") + if __name__ == "__main__": NUM_REQUESTS = 100 # 设置并发请求数 diff --git a/test/test_neo4j.py b/test/test_neo4j.py index 53b967a0..f4c6e1c7 100755 --- a/test/test_neo4j.py +++ b/test/test_neo4j.py @@ -1,8 +1,9 @@ from neo4j import GraphDatabase -from neo4j.exceptions import ServiceUnavailable, AuthError +from neo4j.exceptions import AuthError, ServiceUnavailable # sudo ln -s /snap/core22/1586/usr/sbin/iptables /usr/sbin/iptables + def check_neo4j_status(uri="bolt://localhost:7687", username="neo4j", password="0123456789"): """ 检查 Neo4j 数据库是否可以连接并正常工作。 @@ -30,7 +31,7 @@ def check_neo4j_status(uri="bolt://localhost:7687", username="neo4j", password=" # 确保关闭驱动 driver.close() + # 测试函数 status = check_neo4j_status() print(f"Neo4j status: {status}") -