diff --git a/Makefile b/Makefile index e7feab7f..0a256708 100644 --- a/Makefile +++ b/Makefile @@ -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 \ No newline at end of file + uv run ruff format . + uv run ruff check . --fix \ No newline at end of file diff --git a/scripts/batch_upload.py b/scripts/batch_upload.py index 680dda40..b8ce49ac 100644 --- a/scripts/batch_upload.py +++ b/scripts/batch_upload.py @@ -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]") diff --git a/server/db_manager.py b/server/db_manager.py index b40056a6..d9f1dac7 100644 --- a/server/db_manager.py +++ b/server/db_manager.py @@ -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 diff --git a/server/main.py b/server/main.py index aeb6c31d..a26ead72 100644 --- a/server/main.py +++ b/server/main.py @@ -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 diff --git a/server/models/thread_model.py b/server/models/thread_model.py index 4bdf764b..60d7bde6 100644 --- a/server/models/thread_model.py +++ b/server/models/thread_model.py @@ -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 diff --git a/server/routers/auth_router.py b/server/routers/auth_router.py index 5135ea18..ed1db113 100644 --- a/server/routers/auth_router.py +++ b/server/routers/auth_router.py @@ -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 diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index 9e548476..0b8d16b0 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -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 diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index 168a9158..cf8e16ad 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -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] diff --git a/server/routers/system_router.py b/server/routers/system_router.py index f7945346..0322ad07 100644 --- a/server/routers/system_router.py +++ b/server/routers/system_router.py @@ -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"]) diff --git a/server/utils/auth_middleware.py b/server/utils/auth_middleware.py index 57920f49..ea90159b 100644 --- a/server/utils/auth_middleware.py +++ b/server/utils/auth_middleware.py @@ -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 diff --git a/src/agents/common/mcp.py b/src/agents/common/mcp.py index ae0e5a05..84e88896 100644 --- a/src/agents/common/mcp.py +++ b/src/agents/common/mcp.py @@ -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 diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index b986fabe..87429583 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -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) diff --git a/src/agents/common/utils.py b/src/agents/common/utils.py index c6f70593..84dc5077 100644 --- a/src/agents/common/utils.py +++ b/src/agents/common/utils.py @@ -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 diff --git a/src/agents/react/graph.py b/src/agents/react/graph.py index 08e0a209..774b6a64 100644 --- a/src/agents/react/graph.py +++ b/src/agents/react/graph.py @@ -1,4 +1,3 @@ -import os from pathlib import Path from langchain_core.messages import AnyMessage, SystemMessage diff --git a/src/config/__init__.py b/src/config/__init__.py index 98f1eb6d..c21a897b 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -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" ) diff --git a/src/knowledge/__init__.py b/src/knowledge/__init__.py index b2aa95ee..d2bc9e45 100644 --- a/src/knowledge/__init__.py +++ b/src/knowledge/__init__.py @@ -1 +1,3 @@ from .graphbase import GraphDatabase + +__all__ = ["GraphDatabase"] diff --git a/src/knowledge/chroma_kb.py b/src/knowledge/chroma_kb.py index 17b1a806..bd852c63 100644 --- a/src/knowledge/chroma_kb.py +++ b/src/knowledge/chroma_kb.py @@ -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") diff --git a/src/knowledge/graphbase.py b/src/knowledge/graphbase.py index b80edbb2..bfb163ee 100644 --- a/src/knowledge/graphbase.py +++ b/src/knowledge/graphbase.py @@ -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)" ) # 批量获取嵌入向量 diff --git a/src/knowledge/indexing.py b/src/knowledge/indexing.py index c758e74f..d672847f 100644 --- a/src/knowledge/indexing.py +++ b/src/knowledge/indexing.py @@ -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): diff --git a/src/knowledge/kb_factory.py b/src/knowledge/kb_factory.py index 35549ca3..dfd1ba71 100644 --- a/src/knowledge/kb_factory.py +++ b/src/knowledge/kb_factory.py @@ -1,5 +1,3 @@ -from typing import Any - from src.knowledge.knowledge_base import KBNotFoundError, KnowledgeBase from src.utils import logger diff --git a/src/knowledge/kb_manager.py b/src/knowledge/kb_manager.py index de750d8a..c8b54fce 100644 --- a/src/knowledge/kb_manager.py +++ b/src/knowledge/kb_manager.py @@ -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 diff --git a/src/knowledge/kb_utils.py b/src/knowledge/kb_utils.py index c94449d2..76140a48 100644 --- a/src/knowledge/kb_utils.py +++ b/src/knowledge/kb_utils.py @@ -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]: diff --git a/src/knowledge/knowledge_base.py b/src/knowledge/knowledge_base.py index 7b67cb06..d520d42c 100644 --- a/src/knowledge/knowledge_base.py +++ b/src/knowledge/knowledge_base.py @@ -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 diff --git a/src/knowledge/milvus_kb.py b/src/knowledge/milvus_kb.py index a758f359..59e432ce 100644 --- a/src/knowledge/milvus_kb.py +++ b/src/knowledge/milvus_kb.py @@ -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, diff --git a/src/models/chat_model.py b/src/models/chat_model.py index e70f68b0..2c3cd709 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -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) diff --git a/src/models/embedding.py b/src/models/embedding.py index 170f0c6c..f86d6490 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -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 diff --git a/src/models/rerank_model.py b/src/models/rerank_model.py index 574b0e62..83cb87a3 100644 --- a/src/models/rerank_model.py +++ b/src/models/rerank_model.py @@ -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): diff --git a/src/plugins/_ocr.py b/src/plugins/_ocr.py index 1b387693..8d51568b 100644 --- a/src/plugins/_ocr.py +++ b/src/plugins/_ocr.py @@ -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 = {} diff --git a/src/plugins/mineru.py b/src/plugins/mineru.py index 6acd3cfa..8187ec13 100644 --- a/src/plugins/mineru.py +++ b/src/plugins/mineru.py @@ -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: diff --git a/src/plugins/paddlex.py b/src/plugins/paddlex.py index 5fff01c0..897aaad0 100644 --- a/src/plugins/paddlex.py +++ b/src/plugins/paddlex.py @@ -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 diff --git a/src/utils/logging_config.py b/src/utils/logging_config.py index c3e089c3..3aea635f 100644 --- a/src/utils/logging_config.py +++ b/src/utils/logging_config.py @@ -32,7 +32,12 @@ def setup_logger(name, level="DEBUG", console=True): loguru_logger.add( lambda msg: print(msg, end=""), level=level, - format="{time:MM-DD HH:mm:ss} {level} {name}:{line}: {message}", + format=( + "{time:MM-DD HH:mm:ss} " + "{level} " + "{name}:{line}: " + "{message}" + ), colorize=True, ) diff --git a/src/utils/prompts.py b/src/utils/prompts.py index be740692..e1a32677 100644 --- a/src/utils/prompts.py +++ b/src/utils/prompts.py @@ -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" +) diff --git a/test/test_concurrency.py b/test/test_concurrency.py index 5a78b9a6..72c7a54f 100644 --- a/test/test_concurrency.py +++ b/test/test_concurrency.py @@ -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}" ) # 计算最大并发数