refactor: 清理未使用的导入和优化代码格式

This commit is contained in:
Wenjie Zhang 2025-09-02 01:08:42 +08:00
parent 7641ab3699
commit f1c843addc
33 changed files with 75 additions and 74 deletions

View File

@ -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

View File

@ -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]")

View File

@ -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

View File

@ -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

View File

@ -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

View File

@ -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

View File

@ -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

View File

@ -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]

View File

@ -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"])

View File

@ -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

View File

@ -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

View File

@ -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)

View File

@ -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

View File

@ -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

View File

@ -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"
) )

View File

@ -1 +1,3 @@
from .graphbase import GraphDatabase from .graphbase import GraphDatabase
__all__ = ["GraphDatabase"]

View File

@ -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")

View File

@ -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)"
) )
# 批量获取嵌入向量 # 批量获取嵌入向量

View File

@ -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):

View File

@ -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

View File

@ -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

View File

@ -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]:

View File

@ -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

View File

@ -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,

View File

@ -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)

View File

@ -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

View File

@ -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):

View File

@ -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 = {}

View File

@ -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:

View File

@ -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

View File

@ -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,
) )

View File

@ -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"
)

View File

@ -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}"
) )
# 计算最大并发数 # 计算最大并发数