From 4f8e1e19e1b1e14e928eb398f3b4b1efa6804f35 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sun, 27 Jul 2025 01:02:14 +0800 Subject: [PATCH] =?UTF-8?q?feat(batch=5Fupload):=20=E6=B7=BB=E5=8A=A0?= =?UTF-8?q?=E6=89=B9=E9=87=8F=E4=B8=8A=E4=BC=A0=E5=92=8C=E5=A4=84=E7=90=86?= =?UTF-8?q?=E6=96=87=E4=BB=B6=E5=8A=9F=E8=83=BD=EF=BC=8C=E6=94=AF=E6=8C=81?= =?UTF-8?q?=E6=96=87=E4=BB=B6=E8=AE=B0=E5=BD=95=E7=AE=A1=E7=90=86=EF=BC=8C?= =?UTF-8?q?=E4=BC=98=E5=8C=96API=E4=BA=A4=E4=BA=92=E5=92=8C=E9=94=99?= =?UTF-8?q?=E8=AF=AF=E5=A4=84=E7=90=86=E9=80=BB=E8=BE=91(=E5=B9=B6?= =?UTF-8?q?=E5=8F=91=E8=AE=BF=E9=97=AEMilvus=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- scripts/batch_upload.py | 316 +++++++++++++++++++++++++++++ server/routers/knowledge_router.py | 4 +- src/knowledge/kb_manager.py | 69 +++---- src/knowledge/milvus_kb.py | 73 +++---- 4 files changed, 389 insertions(+), 73 deletions(-) create mode 100644 scripts/batch_upload.py diff --git a/scripts/batch_upload.py b/scripts/batch_upload.py new file mode 100644 index 00000000..bf17ae26 --- /dev/null +++ b/scripts/batch_upload.py @@ -0,0 +1,316 @@ + +import asyncio +import hashlib +import json +import pathlib + +import httpx +import typer +from rich.console import Console +from rich.progress import Progress, SpinnerColumn, BarColumn, TextColumn, TimeElapsedColumn + +app = typer.Typer() +console = Console() + + +async def login(client: httpx.AsyncClient, base_url: str, username: str, password: str) -> str | None: + """Logs in to the API and returns the access token.""" + try: + response = await client.post( + f"{base_url}/auth/token", + data={"username": username, "password": password}, + ) + response.raise_for_status() + return response.json().get("access_token") + except httpx.HTTPStatusError as e: + console.print(f"[bold red]Login failed: {e.response.status_code} - {e.response.text}[/bold red]") + return None + except httpx.RequestError as e: + console.print(f"[bold red]Login request failed: {e}[/bold red]") + return None + + +async def upload_file( + client: httpx.AsyncClient, + base_url: str, + db_id: str, + file_path: pathlib.Path, +) -> str | None: + """Uploads a single file and returns its server-side path.""" + try: + with open(file_path, "rb") as f: + files = {"file": (file_path.name, f, "application/octet-stream")} + response = await client.post( + f"{base_url}/knowledge/files/upload", + params={"db_id": db_id}, + files=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]") + return None + except httpx.RequestError as e: + console.print(f"[bold red]Failed to upload {file_path.name}: {e}[/bold red]") + return None + + +async def process_document( + client: httpx.AsyncClient, + base_url: str, + db_id: str, + server_file_path: str, +) -> bool: + """Triggers the processing of an uploaded file in the knowledge base.""" + try: + response = await client.post( + f"{base_url}/knowledge/databases/{db_id}/documents", + json={"items": [server_file_path], "params": {"content_type": "file"}}, + 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]") + return False + + # Check the specific file's processing status in the items array + items = result.get("items", []) + if not items: + console.print(f"[bold red]No processing result for {server_file_path}[/bold red]") + return False + + # Since we only sent one file, check the first item + item = items[0] + # Check for both 'success' and 'done' status (different APIs might use different status values) + if item.get("status") in ["success", "done"]: + return True + else: + # Get more detailed error information + error_msg = item.get("message", "") + error_detail = item.get("detail", "") + error_reason = item.get("reason", "") + + # Combine all available error information + error_info = [] + if error_msg: + error_info.append(error_msg) + if error_detail: + error_info.append(error_detail) + if error_reason: + error_info.append(error_reason) + + if not error_info: + error_info = ["Unknown error"] + + full_error = " | ".join(error_info) + console.print(f"[bold red]Failed to process {server_file_path}: {full_error}[/bold red]") + + # Also log the full item for debugging + console.print(f"[dim]Debug - Full item response: {item}[/dim]") + 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]") + return False + except httpx.RequestError as e: + console.print(f"[bold red]Failed to process {server_file_path}: {e}[/bold red]") + return False + + +async def worker( + semaphore: asyncio.Semaphore, + client: httpx.AsyncClient, + base_url: str, + db_id: str, + file_path: pathlib.Path, + file_hash: str, + progress: Progress, + upload_task_id: int, + process_task_id: int, +): + """A worker task that uploads and then processes a single file.""" + async with semaphore: + # 1. Upload file + server_file_path = await upload_file(client, base_url, db_id, file_path) + 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 + return file_path, file_hash, "upload_failed" + + # 2. Process file + success = await process_document(client, base_url, db_id, server_file_path) + progress.update(process_task_id, advance=1, postfix=f"Processed {file_path.name}") + + return file_path, file_hash, "success" if success else "processing_failed" + + +def get_file_hash(file_path: pathlib.Path) -> str: + """Calculate SHA256 hash of a file.""" + hash_sha256 = hashlib.sha256() + with open(file_path, "rb") as f: + for chunk in iter(lambda: f.read(4096), b""): + hash_sha256.update(chunk) + return hash_sha256.hexdigest() + + +def load_processed_files(record_file: pathlib.Path) -> set[str]: + """Load the set of processed file hashes from the record file.""" + if not record_file.exists(): + return set() + + try: + with open(record_file) as f: + data = json.load(f) + 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() + + +def save_processed_files(record_file: pathlib.Path, processed_files: set[str]): + """Save the set of processed file hashes to the record file.""" + # Ensure the directory exists + 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) + except OSError as e: + console.print(f"[bold red]Error: Could not save processed files record: {e}[/bold red]") + + +@app.command() +def main( + 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), + 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", 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(4, 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."), +): + """ + Batch upload and process files into a Yuxi-Know knowledge base. + """ + console.print(f"[bold green]Starting batch upload for knowledge base: {db_id}[/bold green]") + + # Load previously processed files + processed_files = load_processed_files(record_file) + console.print(f"Loaded {len(processed_files)} previously processed files from record.") + + # Discover files + glob_method = directory.rglob if recursive else directory.glob + all_files = list(glob_method(pattern)) + if not all_files: + console.print(f"[bold yellow]No files found in '{directory}' matching '{pattern}'. Aborting.[/bold yellow]") + raise typer.Exit() + + # Filter out already processed files + files_to_upload = [] + skipped_files = [] + + for file_path in all_files: + file_hash = get_file_hash(file_path) + if file_hash in processed_files: + skipped_files.append(file_path) + else: + 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]") + raise typer.Exit() + + console.print(f"Found {len(all_files)} total files:") + console.print(f" - [green]New files to process:[/green] {len(files_to_upload)}") + console.print(f" - [blue]Already processed (skipped):[/blue] {len(skipped_files)}") + + async def run(): + async with httpx.AsyncClient() as client: + # Login + token = await login(client, base_url, username, password) + if not token: + raise typer.Exit(code=1) + + client.headers = {"Authorization": f"Bearer {token}"} + + # Setup concurrency and tasks + semaphore = asyncio.Semaphore(concurrency) + tasks = [] + + with Progress( + SpinnerColumn(), + TextColumn("[progress.description]{task.description}"), + BarColumn(), + TextColumn("[progress.percentage]{task.percentage:>3.0f}%"), + TimeElapsedColumn(), + TextColumn("{task.fields[postfix]}"), + console=console, + transient=True, + ) as progress: + upload_task_id = progress.add_task("[bold blue]Uploading...", total=len(files_to_upload), postfix="") + process_task_id = progress.add_task("[bold cyan]Processing...", total=len(files_to_upload), postfix="") + + 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) + ) + tasks.append(task) + + results = await asyncio.gather(*tasks) + + # Summarize results and update processed files record + successful_files = [] + upload_failures = [] + processing_failures = [] + newly_processed_hashes = set() + + for file_path, file_hash, status in results: + if status == 'success': + successful_files.append(file_path) + newly_processed_hashes.add(file_hash) + elif status == 'upload_failed': + upload_failures.append(file_path) + elif status == 'processing_failed': + processing_failures.append(file_path) + # Don't add to processed files if processing failed + + # Save newly processed files to record + 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("[bold green]Batch operation complete.[/bold green]") + console.print(f" - [green]Successful:[/green] {len(successful_files)}") + console.print(f" - [red]Upload Failed:[/red] {len(upload_failures)}") + if upload_failures: + for f in upload_failures: + console.print(f" - {f}") + console.print(f" - [yellow]Processing Failed:[/yellow] {len(processing_failures)}") + if processing_failures: + for f in processing_failures: + console.print(f" - {f}") + + asyncio.run(run()) + +""" +uv run scripts/batch_upload.py \ + --db-id kb_845a9eedb211b349ddb3127ae9be2bfa \ + --directory data.local/XXXX/ \ + --pattern "*.html" \ + --base-url http://172.19.13.5:5050/api \ + --username zwj \ + --password zwj12138 \ + --concurrency 4 \ + --recursive \ + --record-file scripts/tmp/batch_processed_files.txt +""" +if __name__ == "__main__": + app() diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index af216c68..8b899a6f 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -72,7 +72,7 @@ async def update_database_info( """更新知识库信息""" logger.debug(f"Update database {db_id} info: {name}, {description}") try: - database = knowledge_base.update_database(db_id, name, description) + database = await knowledge_base.update_database(db_id, name, description) return {"message": "更新成功", "database": database} except Exception as e: logger.error(f"更新数据库失败 {e}, {traceback.format_exc()}") @@ -83,7 +83,7 @@ async def delete_database(db_id: str, current_user: User = Depends(get_admin_use """删除知识库""" logger.debug(f"Delete database {db_id}") try: - knowledge_base.delete_database(db_id) + await knowledge_base.delete_database(db_id) # 需要重新加载所有智能体,因为工具刷新了 from src.agents import agent_manager diff --git a/src/knowledge/kb_manager.py b/src/knowledge/kb_manager.py index 1ffcebb9..3494d8dd 100644 --- a/src/knowledge/kb_manager.py +++ b/src/knowledge/kb_manager.py @@ -1,6 +1,7 @@ import os import json import time +import asyncio from typing import Any from datetime import datetime @@ -32,6 +33,9 @@ class KnowledgeBaseManager: # 全局数据库元信息 {db_id: metadata_with_kb_type} self.global_databases_meta: dict[str, dict] = {} + # 元数据锁 + self._metadata_lock = asyncio.Lock() + # 加载全局元数据 self._load_global_metadata() @@ -55,16 +59,13 @@ class KnowledgeBaseManager: def _save_global_metadata(self): """保存全局元数据""" meta_file = os.path.join(self.work_dir, "global_metadata.json") - try: - 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) - except Exception as e: - logger.error(f"Failed to save global metadata: {e}") + 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): """初始化已存在的知识库实例""" @@ -139,7 +140,7 @@ class KnowledgeBaseManager: return {"databases": all_databases} - 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: """ 创建数据库 @@ -153,48 +154,44 @@ class KnowledgeBaseManager: Returns: 数据库信息字典 """ - # 验证知识库类型 if not KnowledgeBaseFactory.is_type_supported(kb_type): available_types = list(KnowledgeBaseFactory.get_available_types().keys()) raise ValueError(f"Unsupported knowledge base type: {kb_type}. Available types: {available_types}") - # 获取或创建对应类型的知识库实例 kb_instance = self._get_or_create_kb_instance(kb_type) - # 在知识库实例中创建数据库 db_info = kb_instance.create_database(database_name, description, embed_info, **kwargs) db_id = db_info["db_id"] - # 在全局元数据中记录 - self.global_databases_meta[db_id] = { - "name": database_name, - "description": description, - "kb_type": kb_type, - "created_at": datetime.now().isoformat(), - "additional_params": kwargs.copy() # 将所有额外参数存储在additional_params中 - } - - self._save_global_metadata() + async with self._metadata_lock: + self.global_databases_meta[db_id] = { + "name": database_name, + "description": description, + "kb_type": kb_type, + "created_at": datetime.now().isoformat(), + "additional_params": kwargs.copy() + } + self._save_global_metadata() logger.info(f"Created {kb_type} database: {database_name} ({db_id}) with {kwargs}") return db_info - def delete_database(self, db_id: str) -> dict: + async def delete_database(self, db_id: str) -> dict: """删除数据库""" try: kb_instance = self._get_kb_for_database(db_id) result = kb_instance.delete_database(db_id) - # 从全局元数据中删除 - if db_id in self.global_databases_meta: - del self.global_databases_meta[db_id] - self._save_global_metadata() + async with self._metadata_lock: + if db_id in self.global_databases_meta: + del self.global_databases_meta[db_id] + self._save_global_metadata() return result except KBNotFoundError as e: logger.warning(f"Database {db_id} not found during deletion: {e}") - return {"message": "删除成功"} # 兼容性:即使不存在也返回成功 + return {"message": "删除成功"} async def add_content(self, db_id: str, items: list[str], params: dict | None = None) -> list[dict]: """添加内容(文件/URL)""" @@ -253,16 +250,16 @@ class KnowledgeBaseManager: os.makedirs(general_uploads, exist_ok=True) return general_uploads - def update_database(self, db_id: str, name: str, description: str) -> dict: + async def update_database(self, db_id: str, name: str, description: str) -> dict: """更新数据库""" kb_instance = self._get_kb_for_database(db_id) result = kb_instance.update_database(db_id, name, description) - # 同时更新全局元数据 - if db_id in self.global_databases_meta: - self.global_databases_meta[db_id]["name"] = name - self.global_databases_meta[db_id]["description"] = description - self._save_global_metadata() + async with self._metadata_lock: + if db_id in self.global_databases_meta: + self.global_databases_meta[db_id]["name"] = name + self.global_databases_meta[db_id]["description"] = description + self._save_global_metadata() return result diff --git a/src/knowledge/milvus_kb.py b/src/knowledge/milvus_kb.py index 3d8543b6..1ae5bd3b 100644 --- a/src/knowledge/milvus_kb.py +++ b/src/knowledge/milvus_kb.py @@ -2,6 +2,7 @@ import os import time import traceback import json +import asyncio from pathlib import Path from typing import Any from datetime import datetime @@ -54,6 +55,9 @@ class MilvusKB(KnowledgeBase): self.chunk_size = kwargs.get('chunk_size', 1000) self.chunk_overlap = kwargs.get('chunk_overlap', 200) + # 元数据锁 + self._metadata_lock = asyncio.Lock() + # 初始化连接 self._init_connection() @@ -240,64 +244,59 @@ class MilvusKB(KnowledgeBase): processed_items_info = [] for item in items: - # 准备文件元数据 metadata = prepare_item_metadata(item, content_type, db_id) file_id = metadata["file_id"] filename = metadata["filename"] - # 添加文件记录 file_record = metadata.copy() - del file_record["file_id"] # 从记录中移除file_id,因为它是key - self.files_meta[file_id] = file_record - self._save_metadata() + del file_record["file_id"] + async with self._metadata_lock: + self.files_meta[file_id] = file_record + self._save_metadata() - # 添加 file_id 到返回数据 - file_record = file_record.copy() file_record["file_id"] = file_id try: - # 根据内容类型处理内容 if content_type == "file": markdown_content = await self._process_file_to_markdown(item, params=params) - else: # URL + else: markdown_content = await self._process_url_to_markdown(item, params=params) - # 分割文本成块 chunks = self._split_text_into_chunks(markdown_content, file_id, filename, params) logger.info(f"Split {filename} into {len(chunks)} chunks") - # 准备 Milvus 插入的数据 if chunks: - # 生成嵌入向量 texts = [chunk["content"] for chunk in chunks] embeddings = await embedding_function(texts) - # 准备插入数据 entities = [ - [chunk["id"] for chunk in chunks], # id - [chunk["content"] for chunk in chunks], # content - [chunk["source"] for chunk in chunks], # source - [chunk["chunk_id"] for chunk in chunks], # chunk_id - [chunk["file_id"] for chunk in chunks], # file_id - [chunk["chunk_index"] for chunk in chunks], # chunk_index - embeddings # embedding + [chunk["id"] for chunk in chunks], + [chunk["content"] for chunk in chunks], + [chunk["source"] for chunk in chunks], + [chunk["chunk_id"] for chunk in chunks], + [chunk["file_id"] for chunk in chunks], + [chunk["chunk_index"] for chunk in chunks], + embeddings ] - # 插入到 Milvus - collection.insert(entities) - collection.flush() + def _insert_and_flush(): + collection.insert(entities) + collection.flush() + + await asyncio.to_thread(_insert_and_flush) logger.info(f"Inserted {content_type} {item} into Milvus. Done.") - # 更新状态为完成 - self.files_meta[file_id]["status"] = "done" - self._save_metadata() + async with self._metadata_lock: + self.files_meta[file_id]["status"] = "done" + self._save_metadata() 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() + async with self._metadata_lock: + self.files_meta[file_id]["status"] = "failed" + self._save_metadata() file_record['status'] = "failed" processed_items_info.append(file_record) @@ -367,21 +366,25 @@ class MilvusKB(KnowledgeBase): async def delete_file(self, db_id: str, file_id: str) -> None: """删除文件""" collection = await self._get_milvus_collection(db_id) - if collection: + + def _delete_from_milvus(): + """同步执行 Milvus 删除操作的辅助函数""" try: - # 删除所有相关chunks expr = f'file_id == "{file_id}"' collection.delete(expr) collection.flush() logger.info(f"Deleted chunks for file {file_id} from Milvus") - except Exception as e: logger.error(f"Error deleting file {file_id} from Milvus: {e}") - # 删除文件记录 - if file_id in self.files_meta: - del self.files_meta[file_id] - self._save_metadata() + if collection: + await asyncio.to_thread(_delete_from_milvus) + + # 使用锁确保元数据操作的原子性 + async with self._metadata_lock: + if file_id in self.files_meta: + del self.files_meta[file_id] + self._save_metadata() async def get_file_info(self, db_id: str, file_id: str) -> dict: """获取文件信息和chunks"""