From e45e90fec759b82728c2851748aa1bdce3223c86 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Fri, 1 Aug 2025 16:59:19 +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=E6=96=87=E4=BB=B6=E8=BD=AC=E6=8D=A2=E4=B8=BA?= =?UTF-8?q?Markdown=E5=8A=9F=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增`trans`命令以支持批量将文件转换为Markdown格式 - 实现文件上传和转换的异步处理,支持并发任务 - 处理文件时过滤掉macOS的隐藏文件 - 更新`upload`命令的默认并发数和API基础URL --- scripts/batch_upload.py | 186 +++++++++++++++++++++++++++++++++++++--- 1 file changed, 176 insertions(+), 10 deletions(-) diff --git a/scripts/batch_upload.py b/scripts/batch_upload.py index 1e909663..ad908006 100644 --- a/scripts/batch_upload.py +++ b/scripts/batch_upload.py @@ -211,15 +211,81 @@ def save_processed_files(record_file: pathlib.Path, processed_files: set[str]): console.print(f"[bold red]Error: Could not save processed files record: {e}[/bold red]") +async def convert_to_markdown( + client: httpx.AsyncClient, + base_url: str, + db_id: str, + server_file_path: str, +) -> str | None: + """Calls the file-to-markdown conversion endpoint.""" + try: + response = await client.post( + f"{base_url}/knowledge/files/markdown", + json={"db_id": db_id, "file_path": server_file_path}, + timeout=600, # 10 minutes timeout for conversion + ) + response.raise_for_status() + result = response.json() + if result.get("status") == "success": + return result.get("markdown_content") + else: + 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]") + return None + except httpx.RequestError as e: + console.print(f"[bold red]Request failed for {server_file_path}: {e}[/bold red]") + return None + + +async def trans_worker( + semaphore: asyncio.Semaphore, + client: httpx.AsyncClient, + base_url: str, + db_id: str, + file_path: pathlib.Path, + output_dir: pathlib.Path, + progress: Progress, + task_id: int, +): + """A worker task that uploads a file and converts it to markdown.""" + async with semaphore: + # 1. Upload file + server_file_path = await upload_file(client, base_url, db_id, file_path) + if not server_file_path: + progress.update(task_id, advance=1, postfix=f"[red]Upload failed for {file_path.name}[/red]") + return file_path, "upload_failed" + + # 2. Convert file to markdown + markdown_content = await convert_to_markdown(client, base_url, db_id, server_file_path) + if not markdown_content: + progress.update(task_id, advance=1, postfix=f"[yellow]Conversion failed for {file_path.name}[/yellow]") + return file_path, "conversion_failed" + + # 3. Save markdown content to output directory + 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: + f.write(markdown_content) + progress.update(task_id, advance=1, postfix=f"[green]Converted {file_path.name}[/green]") + return file_path, "success" + except OSError as e: + console.print(f"[bold red]Error saving markdown for {file_path.name}: {e}[/bold red]") + progress.update(task_id, advance=1, postfix=f"[red]Save failed for {file_path.name}[/red]") + return file_path, "save_failed" + + @app.command() -def main( +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), 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."), + 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(4, help="The number of concurrent upload/process tasks."), + 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."), chunk_size: int = typer.Option(1000, help="Chunk size for document processing."), @@ -343,17 +409,117 @@ def main( asyncio.run(run()) + +@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), + 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."), + 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 conversion tasks."), + recursive: bool = typer.Option(False, "--recursive", "-r", help="Search for files recursively in subdirectories."), +): + """ + Batch convert files to Markdown format. + """ + console.print(f"[bold green]Starting batch conversion for files in: {directory}[/bold green]") + output_dir.mkdir(parents=True, exist_ok=True) + + # Discover files + glob_method = directory.rglob if recursive else directory.glob + files_to_convert = list(glob_method(pattern)) + if not files_to_convert: + console.print(f"[bold yellow]No files found in '{directory}' matching '{pattern}'. Aborting.[/bold yellow]") + raise typer.Exit() + + # 过滤掉macos的隐藏文件 + files_to_convert = [f for f in files_to_convert if not f.name.startswith("._")] + + console.print(f"Found {len(files_to_convert)} files to convert.") + + 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: + task_id = progress.add_task("[bold blue]Converting...", total=len(files_to_convert), postfix="") + + 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 + ) + ) + tasks.append(task) + + results = await asyncio.gather(*tasks) + + # Summarize results + successful_files = [] + failed_files = [] + + for file_path, status in results: + if status == 'success': + successful_files.append(file_path) + else: + failed_files.append((file_path, status)) + + console.print("[bold green]Batch conversion complete.[/bold green]") + console.print(f" - [green]Successful:[/green] {len(successful_files)}") + console.print(f" - [red]Failed:[/red] {len(failed_files)}") + if failed_files: + for f, status in failed_files: + console.print(f" - {f} (Reason: {status})") + + asyncio.run(run()) + + """ -uv run scripts/batch_upload.py \ - --db-id kb_845a9eedb211b349ddb3127ae9be2bfa \ - --directory data.local/农业农村局/无锡市农业农村局政府信息公开/ \ +# Example for upload +uv run scripts/batch_upload.py upload \ + --db-id your_kb_id \ + --directory path/to/your/data \ --pattern "*.docx" \ - --base-url http://172.19.13.5:5050/api \ - --username zwj \ - --password zwj12138 \ - --concurrency 1 \ + --base-url http://127.0.0.1:5050/api \ + --username your_username \ + --password your_password \ + --concurrency 4 \ --recursive \ --record-file scripts/tmp/batch_processed_files.txt + +# Example for trans +uv run scripts/batch_upload.py trans \ + --db-id your_kb_id \ + --directory path/to/your/data \ + --output-dir path/to/output_markdown \ + --pattern "*.docx" \ + --base-url http://127.0.0.1:5050/api \ + --username your_username \ + --password your_password \ + --concurrency 4 \ + --recursive """ if __name__ == "__main__": app()