refactor: 清理未使用的导入和优化代码格式
This commit is contained in:
parent
7641ab3699
commit
f1c843addc
4
Makefile
4
Makefile
@ -23,5 +23,5 @@ lint:
|
|||||||
uv run python -m ruff check --select I src
|
uv run python -m ruff check --select I src
|
||||||
|
|
||||||
format format_diff:
|
format format_diff:
|
||||||
uv run ruff format
|
uv run ruff format .
|
||||||
uv run ruff check --select I --fix
|
uv run ruff check . --fix
|
||||||
@ -418,7 +418,8 @@ def upload(
|
|||||||
all_processed_files = processed_files | newly_processed_hashes
|
all_processed_files = processed_files | newly_processed_hashes
|
||||||
save_processed_files(record_file, all_processed_files)
|
save_processed_files(record_file, all_processed_files)
|
||||||
console.print(
|
console.print(
|
||||||
f"[bold green]Updated processed files record with {len(newly_processed_hashes)} new entries.[/bold green]"
|
f"[bold green]Updated processed files record with "
|
||||||
|
f"{len(newly_processed_hashes)} new entries.[/bold green]"
|
||||||
)
|
)
|
||||||
|
|
||||||
console.print("[bold green]Batch operation complete.[/bold green]")
|
console.print("[bold green]Batch operation complete.[/bold green]")
|
||||||
|
|||||||
@ -6,8 +6,6 @@ from sqlalchemy import create_engine
|
|||||||
from sqlalchemy.orm import sessionmaker
|
from sqlalchemy.orm import sessionmaker
|
||||||
|
|
||||||
from server.models import Base
|
from server.models import Base
|
||||||
from server.models.kb_models import KnowledgeDatabase, KnowledgeFile, KnowledgeNode
|
|
||||||
from server.models.thread_model import Thread
|
|
||||||
from server.models.user_model import User
|
from server.models.user_model import User
|
||||||
from src import config
|
from src import config
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
|
|||||||
@ -1,7 +1,6 @@
|
|||||||
import uvicorn
|
import uvicorn
|
||||||
from fastapi import Depends, FastAPI, HTTPException, Request, status
|
from fastapi import FastAPI, Request
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from fastapi.responses import JSONResponse
|
|
||||||
from starlette.middleware.base import BaseHTTPMiddleware
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
|
||||||
from server.routers import router
|
from server.routers import router
|
||||||
|
|||||||
@ -1,6 +1,4 @@
|
|||||||
from sqlalchemy import Column, DateTime, ForeignKey, Integer, String, Text
|
from sqlalchemy import Column, DateTime, Integer, String
|
||||||
from sqlalchemy.dialects.mysql import JSON
|
|
||||||
from sqlalchemy.orm import relationship
|
|
||||||
from sqlalchemy.sql import func
|
from sqlalchemy.sql import func
|
||||||
|
|
||||||
from server.models import Base
|
from server.models import Base
|
||||||
|
|||||||
@ -6,8 +6,8 @@ from pydantic import BaseModel
|
|||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from server.db_manager import db_manager
|
from server.db_manager import db_manager
|
||||||
from server.models.user_model import OperationLog, User
|
from server.models.user_model import User
|
||||||
from server.utils.auth_middleware import get_admin_user, get_current_user, get_db, get_superadmin_user, oauth2_scheme
|
from server.utils.auth_middleware import get_admin_user, get_current_user, get_db
|
||||||
from server.utils.auth_utils import AuthUtils
|
from server.utils.auth_utils import AuthUtils
|
||||||
from server.utils.common_utils import log_operation
|
from server.utils.common_utils import log_operation
|
||||||
|
|
||||||
|
|||||||
@ -1,11 +1,9 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import os
|
|
||||||
import time
|
|
||||||
import traceback
|
import traceback
|
||||||
import uuid
|
import uuid
|
||||||
|
|
||||||
from fastapi import APIRouter, Body, Depends, HTTPException, Query
|
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||||
from fastapi.responses import StreamingResponse
|
from fastapi.responses import StreamingResponse
|
||||||
from langchain_core.messages import AIMessageChunk, HumanMessage
|
from langchain_core.messages import AIMessageChunk, HumanMessage
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|||||||
@ -1,13 +1,12 @@
|
|||||||
import asyncio
|
|
||||||
import os
|
import os
|
||||||
import traceback
|
import traceback
|
||||||
|
|
||||||
from fastapi import APIRouter, Body, Depends, File, Form, HTTPException, Query, UploadFile
|
from fastapi import APIRouter, Body, Depends, File, HTTPException, Query, UploadFile
|
||||||
from fastapi.responses import FileResponse
|
from fastapi.responses import FileResponse
|
||||||
|
|
||||||
from server.models.user_model import User
|
from server.models.user_model import User
|
||||||
from server.utils.auth_middleware import get_admin_user
|
from server.utils.auth_middleware import get_admin_user
|
||||||
from src import config, executor, knowledge_base
|
from src import config, knowledge_base
|
||||||
from src.knowledge.indexing import process_file_to_markdown
|
from src.knowledge.indexing import process_file_to_markdown
|
||||||
from src.utils import hashstr, logger
|
from src.utils import hashstr, logger
|
||||||
|
|
||||||
@ -41,7 +40,8 @@ async def create_database(
|
|||||||
):
|
):
|
||||||
"""创建知识库"""
|
"""创建知识库"""
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Create database {database_name} with kb_type {kb_type}, additional_params {additional_params}, llm_info {llm_info}"
|
f"Create database {database_name} with kb_type {kb_type}, "
|
||||||
|
f"additional_params {additional_params}, llm_info {llm_info}"
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
embed_info = config.embed_model_names[embed_model_name]
|
embed_info = config.embed_model_names[embed_model_name]
|
||||||
|
|||||||
@ -4,11 +4,11 @@ from pathlib import Path
|
|||||||
|
|
||||||
import requests
|
import requests
|
||||||
import yaml
|
import yaml
|
||||||
from fastapi import APIRouter, Body, Depends, HTTPException, Request
|
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||||
|
|
||||||
from server.models.user_model import User
|
from server.models.user_model import User
|
||||||
from server.utils.auth_middleware import get_admin_user, get_superadmin_user
|
from server.utils.auth_middleware import get_admin_user, get_superadmin_user
|
||||||
from src import config, graph_base, knowledge_base
|
from src import config, graph_base
|
||||||
from src.utils.logging_config import logger
|
from src.utils.logging_config import logger
|
||||||
|
|
||||||
system = APIRouter(prefix="/system", tags=["system"])
|
system = APIRouter(prefix="/system", tags=["system"])
|
||||||
|
|||||||
@ -2,7 +2,7 @@ import re
|
|||||||
|
|
||||||
from fastapi import Depends, HTTPException, status
|
from fastapi import Depends, HTTPException, status
|
||||||
from fastapi.security import OAuth2PasswordBearer
|
from fastapi.security import OAuth2PasswordBearer
|
||||||
from jose import JWTError, jwt
|
from jose import JWTError
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from server.db_manager import db_manager
|
from server.db_manager import db_manager
|
||||||
|
|||||||
@ -1,13 +1,11 @@
|
|||||||
"""MCP Client setup and management for LangGraph ReAct Agent."""
|
"""MCP Client setup and management for LangGraph ReAct Agent."""
|
||||||
|
|
||||||
import traceback
|
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
|
|
||||||
from langchain_mcp_adapters.client import ( # type: ignore[import-untyped]
|
from langchain_mcp_adapters.client import ( # type: ignore[import-untyped]
|
||||||
MultiServerMCPClient,
|
MultiServerMCPClient,
|
||||||
)
|
)
|
||||||
from langchain_mcp_adapters.tools import load_mcp_tools
|
|
||||||
|
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
|
|
||||||
|
|||||||
@ -1,5 +1,4 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import inspect
|
|
||||||
import traceback
|
import traceback
|
||||||
from typing import Annotated, Any
|
from typing import Annotated, Any
|
||||||
|
|
||||||
@ -18,7 +17,8 @@ def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge
|
|||||||
logger.debug(f"Querying knowledge graph with: {query}")
|
logger.debug(f"Querying knowledge graph with: {query}")
|
||||||
result = graph_base.query_node(query, hops=2, return_format="triples")
|
result = graph_base.query_node(query, hops=2, return_format="triples")
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Knowledge graph query returned {len(result.get('triples', [])) if isinstance(result, dict) else 'N/A'} triples"
|
f"Knowledge graph query returned "
|
||||||
|
f"{len(result.get('triples', [])) if isinstance(result, dict) else 'N/A'} triples"
|
||||||
)
|
)
|
||||||
return result
|
return result
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@ -79,7 +79,10 @@ def get_kb_based_tools() -> list:
|
|||||||
tool_id = f"query_{db_id[:8]}"
|
tool_id = f"query_{db_id[:8]}"
|
||||||
|
|
||||||
# 构建工具描述
|
# 构建工具描述
|
||||||
description = f"使用 {retrieve_info['name']} 知识库进行检索。\n下面是这个知识库的描述:\n{retrieve_info['description'] or '没有描述。'} "
|
description = (
|
||||||
|
f"使用 {retrieve_info['name']} 知识库进行检索。\n"
|
||||||
|
f"下面是这个知识库的描述:\n{retrieve_info['description'] or '没有描述。'} "
|
||||||
|
)
|
||||||
|
|
||||||
# 使用工厂函数创建检索器包装函数,避免闭包问题
|
# 使用工厂函数创建检索器包装函数,避免闭包问题
|
||||||
retriever_wrapper = _create_retriever_wrapper(db_id, retrieve_info)
|
retriever_wrapper = _create_retriever_wrapper(db_id, retrieve_info)
|
||||||
|
|||||||
@ -1,7 +1,6 @@
|
|||||||
import asyncio
|
|
||||||
import os
|
import os
|
||||||
import traceback
|
import traceback
|
||||||
from datetime import UTC, datetime, timezone
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
from langchain_core.language_models import BaseChatModel
|
from langchain_core.language_models import BaseChatModel
|
||||||
from langchain_core.messages import AIMessageChunk, ToolMessage
|
from langchain_core.messages import AIMessageChunk, ToolMessage
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
import os
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from langchain_core.messages import AnyMessage, SystemMessage
|
from langchain_core.messages import AnyMessage, SystemMessage
|
||||||
|
|||||||
@ -52,8 +52,10 @@ class Config(SimpleConfig):
|
|||||||
self.add_item(
|
self.add_item(
|
||||||
"enable_web_search",
|
"enable_web_search",
|
||||||
default=False,
|
default=False,
|
||||||
des="是否开启网页搜索(注:现阶段会根据 TAVILY_API_KEY 自动开启,无法手动配置,将会在下个版本移除此配置项)",
|
des=(
|
||||||
) # noqa: E501
|
"是否开启网页搜索(注:现阶段会根据 TAVILY_API_KEY 自动开启,无法手动配置,将会在下个版本移除此配置项)"
|
||||||
|
),
|
||||||
|
)
|
||||||
# 默认智能体配置
|
# 默认智能体配置
|
||||||
self.add_item("default_agent_id", default="", des="默认智能体ID")
|
self.add_item("default_agent_id", default="", des="默认智能体ID")
|
||||||
# 模型配置
|
# 模型配置
|
||||||
@ -133,11 +135,14 @@ class Config(SimpleConfig):
|
|||||||
if self.model_dir:
|
if self.model_dir:
|
||||||
if os.path.exists(self.model_dir):
|
if os.path.exists(self.model_dir):
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"The model directory ({self.model_dir}) contains the following folders: {os.listdir(self.model_dir)}"
|
f"The model directory ({self.model_dir}) "
|
||||||
|
f"contains the following folders: {os.listdir(self.model_dir)}"
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.warning(
|
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;"
|
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"
|
"For example, the mapping in the docker-compose file"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -1 +1,3 @@
|
|||||||
from .graphbase import GraphDatabase
|
from .graphbase import GraphDatabase
|
||||||
|
|
||||||
|
__all__ = ["GraphDatabase"]
|
||||||
|
|||||||
@ -1,17 +1,12 @@
|
|||||||
import json
|
|
||||||
import os
|
import os
|
||||||
import time
|
|
||||||
import traceback
|
import traceback
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from typing import Any
|
||||||
from typing import Any, Optional
|
|
||||||
|
|
||||||
import chromadb
|
import chromadb
|
||||||
from chromadb.api.types import Documents, EmbeddingFunction, Embeddings
|
|
||||||
from chromadb.config import Settings
|
from chromadb.config import Settings
|
||||||
from chromadb.utils.embedding_functions import OpenAIEmbeddingFunction
|
from chromadb.utils.embedding_functions import OpenAIEmbeddingFunction
|
||||||
|
|
||||||
from src import config
|
|
||||||
from src.knowledge.indexing import process_file_to_markdown, process_url_to_markdown
|
from src.knowledge.indexing import process_file_to_markdown, process_url_to_markdown
|
||||||
from src.knowledge.kb_utils import (
|
from src.knowledge.kb_utils import (
|
||||||
get_embedding_config,
|
get_embedding_config,
|
||||||
@ -20,7 +15,7 @@ from src.knowledge.kb_utils import (
|
|||||||
split_text_into_qa_chunks,
|
split_text_into_qa_chunks,
|
||||||
)
|
)
|
||||||
from src.knowledge.knowledge_base import KnowledgeBase
|
from src.knowledge.knowledge_base import KnowledgeBase
|
||||||
from src.utils import hashstr, logger
|
from src.utils import logger
|
||||||
|
|
||||||
|
|
||||||
class ChromaKB(KnowledgeBase):
|
class ChromaKB(KnowledgeBase):
|
||||||
@ -84,7 +79,8 @@ class ChromaKB(KnowledgeBase):
|
|||||||
# 如果模型不匹配,删除现有集合并重新创建
|
# 如果模型不匹配,删除现有集合并重新创建
|
||||||
if current_model != expected_model:
|
if current_model != expected_model:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Collection {collection_name} uses model '{current_model}', but expected '{expected_model}'. Recreating collection."
|
f"Collection {collection_name} uses model '{current_model}', "
|
||||||
|
f"but expected '{expected_model}'. Recreating collection."
|
||||||
)
|
)
|
||||||
self.chroma_client.delete_collection(name=collection_name)
|
self.chroma_client.delete_collection(name=collection_name)
|
||||||
raise Exception("Model mismatch, recreating collection")
|
raise Exception("Model mismatch, recreating collection")
|
||||||
|
|||||||
@ -4,7 +4,6 @@ import traceback
|
|||||||
import warnings
|
import warnings
|
||||||
|
|
||||||
from neo4j import GraphDatabase as GD
|
from neo4j import GraphDatabase as GD
|
||||||
from neo4j import Query
|
|
||||||
|
|
||||||
from src import config
|
from src import config
|
||||||
from src.models import select_embedding_model
|
from src.models import select_embedding_model
|
||||||
@ -301,7 +300,8 @@ class GraphDatabase:
|
|||||||
for i in range(0, total_entities, max_batch_size):
|
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(
|
logger.debug(
|
||||||
f"Processing entities batch {i // max_batch_size + 1}/{(total_entities - 1) // max_batch_size + 1} ({len(batch_entities)} entities)"
|
f"Processing entities batch {i // max_batch_size + 1}/"
|
||||||
|
f"{(total_entities - 1) // max_batch_size + 1} ({len(batch_entities)} entities)"
|
||||||
)
|
)
|
||||||
|
|
||||||
# 批量获取嵌入向量
|
# 批量获取嵌入向量
|
||||||
|
|||||||
@ -2,7 +2,6 @@ import asyncio
|
|||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from langchain.schema.document import Document
|
|
||||||
from langchain.text_splitter import RecursiveCharacterTextSplitter
|
from langchain.text_splitter import RecursiveCharacterTextSplitter
|
||||||
from langchain_community.document_loaders import (
|
from langchain_community.document_loaders import (
|
||||||
CSVLoader,
|
CSVLoader,
|
||||||
@ -14,7 +13,7 @@ from langchain_community.document_loaders import (
|
|||||||
UnstructuredMarkdownLoader,
|
UnstructuredMarkdownLoader,
|
||||||
)
|
)
|
||||||
|
|
||||||
from src.utils import hashstr, logger
|
from src.utils import logger
|
||||||
|
|
||||||
|
|
||||||
def chunk_with_parser(file_path, params=None):
|
def chunk_with_parser(file_path, params=None):
|
||||||
|
|||||||
@ -1,5 +1,3 @@
|
|||||||
from typing import Any
|
|
||||||
|
|
||||||
from src.knowledge.knowledge_base import KBNotFoundError, KnowledgeBase
|
from src.knowledge.knowledge_base import KBNotFoundError, KnowledgeBase
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
|
|
||||||
|
|||||||
@ -1,12 +1,10 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
import time
|
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from src.knowledge.kb_factory import KnowledgeBaseFactory
|
from src.knowledge.kb_factory import KnowledgeBaseFactory
|
||||||
from src.knowledge.knowledge_base import KBNotFoundError, KBOperationError, KnowledgeBase
|
from src.knowledge.knowledge_base import KBNotFoundError, KnowledgeBase
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -1,12 +1,11 @@
|
|||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from langchain_text_splitters import MarkdownTextSplitter
|
from langchain_text_splitters import MarkdownTextSplitter
|
||||||
|
|
||||||
from src import config
|
from src import config
|
||||||
from src.utils import get_docker_safe_url, hashstr, logger
|
from src.utils import hashstr, logger
|
||||||
|
|
||||||
|
|
||||||
def split_text_into_chunks(text: str, file_id: str, filename: str, params: dict = {}) -> list[dict]:
|
def split_text_into_chunks(text: str, file_id: str, filename: str, params: dict = {}) -> list[dict]:
|
||||||
|
|||||||
@ -2,9 +2,7 @@ import json
|
|||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from collections.abc import AsyncGenerator
|
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
|
|||||||
@ -1,16 +1,11 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
|
||||||
import os
|
import os
|
||||||
import time
|
|
||||||
import traceback
|
import traceback
|
||||||
from datetime import datetime
|
|
||||||
from functools import partial
|
from functools import partial
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from pymilvus import Collection, CollectionSchema, DataType, FieldSchema, connections, db, utility
|
from pymilvus import Collection, CollectionSchema, DataType, FieldSchema, connections, db, utility
|
||||||
|
|
||||||
from src import config
|
|
||||||
from src.knowledge.indexing import process_file_to_markdown, process_url_to_markdown
|
from src.knowledge.indexing import process_file_to_markdown, process_url_to_markdown
|
||||||
from src.knowledge.kb_utils import (
|
from src.knowledge.kb_utils import (
|
||||||
get_embedding_config,
|
get_embedding_config,
|
||||||
|
|||||||
@ -1,7 +1,5 @@
|
|||||||
import os
|
import os
|
||||||
|
|
||||||
import requests
|
|
||||||
from langchain_openai import ChatOpenAI
|
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
|
|
||||||
from src.utils import get_docker_safe_url, logger
|
from src.utils import get_docker_safe_url, logger
|
||||||
@ -38,7 +36,10 @@ class OpenAIBase:
|
|||||||
yield chunk.choices[0].delta
|
yield chunk.choices[0].delta
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
err = f"Error streaming response: {e}, URL: {self.base_url}, API Key: {self.api_key[:5]}***, Model: {self.model_name}"
|
err = (
|
||||||
|
f"Error streaming response: {e}, URL: {self.base_url}, "
|
||||||
|
f"API Key: {self.api_key[:5]}***, Model: {self.model_name}"
|
||||||
|
)
|
||||||
logger.error(err)
|
logger.error(err)
|
||||||
raise Exception(err)
|
raise Exception(err)
|
||||||
|
|
||||||
|
|||||||
@ -6,7 +6,6 @@ from abc import ABC, abstractmethod
|
|||||||
import httpx
|
import httpx
|
||||||
import requests
|
import requests
|
||||||
|
|
||||||
from src import config
|
|
||||||
from src.utils import get_docker_safe_url, hashstr, logger
|
from src.utils import get_docker_safe_url, hashstr, logger
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -5,7 +5,7 @@ import numpy as np
|
|||||||
import requests
|
import requests
|
||||||
|
|
||||||
from src import config
|
from src import config
|
||||||
from src.utils import get_docker_safe_url, logger
|
from src.utils import get_docker_safe_url
|
||||||
|
|
||||||
|
|
||||||
def sigmoid(x):
|
def sigmoid(x):
|
||||||
|
|||||||
@ -3,7 +3,6 @@ import time
|
|||||||
import uuid
|
import uuid
|
||||||
from argparse import ArgumentParser
|
from argparse import ArgumentParser
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
import fitz # fitz就是pip install PyMuPDF
|
import fitz # fitz就是pip install PyMuPDF
|
||||||
import numpy as np # Added import for numpy
|
import numpy as np # Added import for numpy
|
||||||
@ -11,7 +10,7 @@ from PIL import Image
|
|||||||
from rapidocr_onnxruntime import RapidOCR
|
from rapidocr_onnxruntime import RapidOCR
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
|
||||||
from src.utils import is_text_pdf, logger
|
from src.utils import logger
|
||||||
|
|
||||||
GOLBAL_STATE = {}
|
GOLBAL_STATE = {}
|
||||||
|
|
||||||
|
|||||||
@ -196,7 +196,10 @@ def parse_doc(
|
|||||||
Parameter description:
|
Parameter description:
|
||||||
path_list: List of document paths to be parsed, can be PDF or image files.
|
path_list: List of document paths to be parsed, can be PDF or image files.
|
||||||
output_dir: Output directory for storing parsing results.
|
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']。
|
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.
|
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"
|
Adapted only for the case where the backend is set to "pipeline"
|
||||||
backend: the backend for parsing pdf:
|
backend: the backend for parsing pdf:
|
||||||
|
|||||||
@ -3,7 +3,7 @@ import json
|
|||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Optional
|
from typing import Any
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
|
|
||||||
|
|||||||
@ -32,7 +32,12 @@ def setup_logger(name, level="DEBUG", console=True):
|
|||||||
loguru_logger.add(
|
loguru_logger.add(
|
||||||
lambda msg: print(msg, end=""),
|
lambda msg: print(msg, end=""),
|
||||||
level=level,
|
level=level,
|
||||||
format="<green>{time:MM-DD HH:mm:ss}</green> <level>{level}</level> <cyan>{name}:{line}</cyan>: <level>{message}</level>",
|
format=(
|
||||||
|
"<green>{time:MM-DD HH:mm:ss}</green> "
|
||||||
|
"<level>{level}</level> "
|
||||||
|
"<cyan>{name}:{line}</cyan>: "
|
||||||
|
"<level>{message}</level>"
|
||||||
|
),
|
||||||
colorize=True,
|
colorize=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -66,4 +66,8 @@ keywords_prompt_template = """
|
|||||||
<文本>{text}</文本>
|
<文本>{text}</文本>
|
||||||
"""
|
"""
|
||||||
|
|
||||||
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"
|
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\nPassage:\n"
|
||||||
|
)
|
||||||
|
|||||||
@ -18,7 +18,9 @@ async def make_request(session: aiohttp.ClientSession, request_id: int) -> dict:
|
|||||||
end_time = time.time()
|
end_time = time.time()
|
||||||
duration = end_time - start_time
|
duration = end_time - start_time
|
||||||
print(
|
print(
|
||||||
f"请求 {request_id} 完成时间: {time.strftime('%H:%M:%S', time.localtime(end_time))} (耗时: {duration:.2f}秒)"
|
f"请求 {request_id} 完成时间: "
|
||||||
|
f"{time.strftime('%H:%M:%S', time.localtime(end_time))} "
|
||||||
|
f"(耗时: {duration:.2f}秒)"
|
||||||
)
|
)
|
||||||
return {
|
return {
|
||||||
"request_id": request_id,
|
"request_id": request_id,
|
||||||
@ -32,7 +34,9 @@ async def make_request(session: aiohttp.ClientSession, request_id: int) -> dict:
|
|||||||
end_time = time.time()
|
end_time = time.time()
|
||||||
duration = end_time - start_time
|
duration = end_time - start_time
|
||||||
print(
|
print(
|
||||||
f"请求 {request_id} 失败时间: {time.strftime('%H:%M:%S', time.localtime(end_time))} (耗时: {duration:.2f}秒)"
|
f"请求 {request_id} 失败时间: "
|
||||||
|
f"{time.strftime('%H:%M:%S', time.localtime(end_time))} "
|
||||||
|
f"(耗时: {duration:.2f}秒)"
|
||||||
)
|
)
|
||||||
return {
|
return {
|
||||||
"request_id": request_id,
|
"request_id": request_id,
|
||||||
@ -102,7 +106,10 @@ def analyze_results(results: list[dict]) -> None:
|
|||||||
active_requests.append(end_time)
|
active_requests.append(end_time)
|
||||||
|
|
||||||
print(
|
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}"
|
f"{result['request_id']:^7} "
|
||||||
|
f"{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}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# 计算最大并发数
|
# 计算最大并发数
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user