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
format format_diff:
uv run ruff format
uv run ruff check --select I --fix
uv run ruff format .
uv run ruff check . --fix

View File

@ -418,7 +418,8 @@ def upload(
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]"
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]")

View File

@ -6,8 +6,6 @@ from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
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 src import config
from src.utils import logger

View File

@ -1,7 +1,6 @@
import uvicorn
from fastapi import Depends, FastAPI, HTTPException, Request, status
from fastapi import FastAPI, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from starlette.middleware.base import BaseHTTPMiddleware
from server.routers import router

View File

@ -1,6 +1,4 @@
from sqlalchemy import Column, DateTime, ForeignKey, Integer, String, Text
from sqlalchemy.dialects.mysql import JSON
from sqlalchemy.orm import relationship
from sqlalchemy import Column, DateTime, Integer, String
from sqlalchemy.sql import func
from server.models import Base

View File

@ -6,8 +6,8 @@ from pydantic import BaseModel
from sqlalchemy.orm import Session
from server.db_manager import db_manager
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.models.user_model import User
from server.utils.auth_middleware import get_admin_user, get_current_user, get_db
from server.utils.auth_utils import AuthUtils
from server.utils.common_utils import log_operation

View File

@ -1,11 +1,9 @@
import asyncio
import json
import os
import time
import traceback
import uuid
from fastapi import APIRouter, Body, Depends, HTTPException, Query
from fastapi import APIRouter, Body, Depends, HTTPException
from fastapi.responses import StreamingResponse
from langchain_core.messages import AIMessageChunk, HumanMessage
from pydantic import BaseModel

View File

@ -1,13 +1,12 @@
import asyncio
import os
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 server.models.user_model import 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.utils import hashstr, logger
@ -41,7 +40,8 @@ async def create_database(
):
"""创建知识库"""
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:
embed_info = config.embed_model_names[embed_model_name]

View File

@ -4,11 +4,11 @@ from pathlib import Path
import requests
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.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
system = APIRouter(prefix="/system", tags=["system"])

View File

@ -2,7 +2,7 @@ import re
from fastapi import Depends, HTTPException, status
from fastapi.security import OAuth2PasswordBearer
from jose import JWTError, jwt
from jose import JWTError
from sqlalchemy.orm import Session
from server.db_manager import db_manager

View File

@ -1,13 +1,11 @@
"""MCP Client setup and management for LangGraph ReAct Agent."""
import traceback
from collections.abc import Callable
from typing import Any, cast
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

View File

@ -1,5 +1,4 @@
import asyncio
import inspect
import traceback
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}")
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"
f"Knowledge graph query returned "
f"{len(result.get('triples', [])) if isinstance(result, dict) else 'N/A'} triples"
)
return result
except Exception as e:
@ -79,7 +79,10 @@ def get_kb_based_tools() -> list:
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)

View File

@ -1,7 +1,6 @@
import asyncio
import os
import traceback
from datetime import UTC, datetime, timezone
from datetime import UTC, datetime
from langchain_core.language_models import BaseChatModel
from langchain_core.messages import AIMessageChunk, ToolMessage

View File

@ -1,4 +1,3 @@
import os
from pathlib import Path
from langchain_core.messages import AnyMessage, SystemMessage

View File

@ -52,8 +52,10 @@ class Config(SimpleConfig):
self.add_item(
"enable_web_search",
default=False,
des="是否开启网页搜索(注:现阶段会根据 TAVILY_API_KEY 自动开启,无法手动配置,将会在下个版本移除此配置项)",
) # noqa: E501
des=(
"是否开启网页搜索(注:现阶段会根据 TAVILY_API_KEY 自动开启,无法手动配置,将会在下个版本移除此配置项)"
),
)
# 默认智能体配置
self.add_item("default_agent_id", default="", des="默认智能体ID")
# 模型配置
@ -133,11 +135,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)}"
f"The model directory {self.model_dir} "
f"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;"
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"
)

View File

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

View File

@ -1,17 +1,12 @@
import json
import os
import time
import traceback
from datetime import datetime
from pathlib import Path
from typing import Any, Optional
from typing import Any
import chromadb
from chromadb.api.types import Documents, EmbeddingFunction, Embeddings
from chromadb.config import Settings
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.kb_utils import (
get_embedding_config,
@ -20,7 +15,7 @@ from src.knowledge.kb_utils import (
split_text_into_qa_chunks,
)
from src.knowledge.knowledge_base import KnowledgeBase
from src.utils import hashstr, logger
from src.utils import logger
class ChromaKB(KnowledgeBase):
@ -84,7 +79,8 @@ class ChromaKB(KnowledgeBase):
# 如果模型不匹配,删除现有集合并重新创建
if current_model != expected_model:
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)
raise Exception("Model mismatch, recreating collection")

View File

@ -4,7 +4,6 @@ import traceback
import warnings
from neo4j import GraphDatabase as GD
from neo4j import Query
from src import config
from src.models import select_embedding_model
@ -301,7 +300,8 @@ class GraphDatabase:
for i in range(0, total_entities, max_batch_size):
batch_entities = nodes_without_embedding[i : i + max_batch_size]
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
from pathlib import Path
from langchain.schema.document import Document
from langchain.text_splitter import RecursiveCharacterTextSplitter
from langchain_community.document_loaders import (
CSVLoader,
@ -14,7 +13,7 @@ from langchain_community.document_loaders import (
UnstructuredMarkdownLoader,
)
from src.utils import hashstr, logger
from src.utils import logger
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.utils import logger

View File

@ -1,12 +1,10 @@
import asyncio
import json
import os
import time
from datetime import datetime
from typing import Any
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

View File

@ -1,12 +1,11 @@
import os
import time
from pathlib import Path
from typing import Any
from langchain_text_splitters import MarkdownTextSplitter
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]:

View File

@ -2,9 +2,7 @@ import json
import os
import time
from abc import ABC, abstractmethod
from collections.abc import AsyncGenerator
from datetime import datetime
from pathlib import Path
from typing import Any
from src.utils import logger

View File

@ -1,16 +1,11 @@
import asyncio
import json
import os
import time
import traceback
from datetime import datetime
from functools import partial
from pathlib import Path
from typing import Any
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.kb_utils import (
get_embedding_config,

View File

@ -1,7 +1,5 @@
import os
import requests
from langchain_openai import ChatOpenAI
from openai import OpenAI
from src.utils import get_docker_safe_url, logger
@ -38,7 +36,10 @@ class OpenAIBase:
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}"
err = (
f"Error streaming response: {e}, URL: {self.base_url}, "
f"API Key: {self.api_key[:5]}***, Model: {self.model_name}"
)
logger.error(err)
raise Exception(err)

View File

@ -6,7 +6,6 @@ from abc import ABC, abstractmethod
import httpx
import requests
from src import config
from src.utils import get_docker_safe_url, hashstr, logger

View File

@ -5,7 +5,7 @@ import numpy as np
import requests
from src import config
from src.utils import get_docker_safe_url, logger
from src.utils import get_docker_safe_url
def sigmoid(x):

View File

@ -3,7 +3,6 @@ import time
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
@ -11,7 +10,7 @@ from PIL import Image
from rapidocr_onnxruntime import RapidOCR
from tqdm import tqdm
from src.utils import is_text_pdf, logger
from src.utils import logger
GOLBAL_STATE = {}

View File

@ -196,7 +196,10 @@ def parse_doc(
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']
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:

View File

@ -3,7 +3,7 @@ import json
import os
import time
from pathlib import Path
from typing import Any, Optional
from typing import Any
import requests

View File

@ -32,7 +32,12 @@ def setup_logger(name, level="DEBUG", console=True):
loguru_logger.add(
lambda msg: print(msg, end=""),
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,
)

View File

@ -66,4 +66,8 @@ keywords_prompt_template = """
<文本>{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()
duration = end_time - start_time
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 {
"request_id": request_id,
@ -32,7 +34,9 @@ async def make_request(session: aiohttp.ClientSession, request_id: int) -> dict:
end_time = time.time()
duration = end_time - start_time
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 {
"request_id": request_id,
@ -102,7 +106,10 @@ def analyze_results(results: list[dict]) -> None:
active_requests.append(end_time)
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}"
)
# 计算最大并发数