diff --git a/src/config/app.py b/src/config/app.py index b92649b0..9624011b 100644 --- a/src/config/app.py +++ b/src/config/app.py @@ -171,9 +171,7 @@ class Config(SimpleConfig): self.enable_web_search = True self.valuable_model_provider = [k for k, v in self.model_provider_status.items() if v] - assert len(self.valuable_model_provider) > 0, ( - "No model provider available, please check your `.env` file." - ) + assert len(self.valuable_model_provider) > 0, "No model provider available, please check your `.env` file." def load(self): """根据传入的文件覆盖掉默认配置""" @@ -223,5 +221,4 @@ class Config(SimpleConfig): return json.loads(str(self)) - config = Config() diff --git a/src/models/chat.py b/src/models/chat.py index 2bf3478e..aa06c7a3 100644 --- a/src/models/chat.py +++ b/src/models/chat.py @@ -5,7 +5,7 @@ from openai import OpenAI from tenacity import before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential from src import config -from src.utils import get_docker_safe_url, logger +from src.utils import logger class OpenAIBase: @@ -79,8 +79,6 @@ class OpenModel(OpenAIBase): super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) - - class GeneralResponse: def __init__(self, content): self.content = content @@ -110,8 +108,6 @@ def select_model(model_provider, model_name=None): raise ValueError(f"Model provider {model_provider} load failed, {e} \n {traceback.format_exc()}") - - async def test_chat_model_status(provider: str, model_name: str) -> dict: """ 测试指定聊天模型的状态 diff --git a/test/api/test_graph_router.py b/test/api/test_graph_router.py index 34b492db..ccfb8640 100644 --- a/test/api/test_graph_router.py +++ b/test/api/test_graph_router.py @@ -15,9 +15,7 @@ async def test_graph_routes_require_auth(test_client): async def test_standard_user_cannot_access_graph_endpoints(test_client, standard_user): - response = await test_client.get( - "/api/graph/lightrag/databases", headers=standard_user["headers"] - ) + response = await test_client.get("/api/graph/lightrag/databases", headers=standard_user["headers"]) assert response.status_code == 403 diff --git a/test/api/test_knowledge_router.py b/test/api/test_knowledge_router.py index 932153c4..ec95b517 100644 --- a/test/api/test_knowledge_router.py +++ b/test/api/test_knowledge_router.py @@ -47,7 +47,5 @@ async def test_knowledge_routes_enforce_permissions(test_client, standard_user, forbidden_list = await test_client.get("/api/knowledge/databases", headers=standard_user["headers"]) assert forbidden_list.status_code == 403 - forbidden_get = await test_client.get( - f"/api/knowledge/databases/{db_id}", headers=standard_user["headers"] - ) + forbidden_get = await test_client.get(f"/api/knowledge/databases/{db_id}", headers=standard_user["headers"]) assert forbidden_get.status_code == 403 diff --git a/test/api/test_system_router.py b/test/api/test_system_router.py index ae7d19db..65843d23 100644 --- a/test/api/test_system_router.py +++ b/test/api/test_system_router.py @@ -25,9 +25,7 @@ async def test_info_endpoint_is_public(test_client): async def test_config_endpoints_require_admin(test_client, standard_user): assert (await test_client.get("/api/system/config")).status_code == 401 - assert ( - await test_client.get("/api/system/config", headers=standard_user["headers"]) - ).status_code == 403 + assert (await test_client.get("/api/system/config", headers=standard_user["headers"])).status_code == 403 async def test_admin_can_fetch_config_and_reload_info(test_client, admin_headers): diff --git a/test/conftest.py b/test/conftest.py index e160bca0..a7b41f89 100644 --- a/test/conftest.py +++ b/test/conftest.py @@ -139,7 +139,9 @@ async def knowledge_database(test_client: httpx.AsyncClient, admin_headers: dict headers=admin_headers, ) if create_response.status_code != 200: - pytest.fail(f"Failed to create knowledge database (status={create_response.status_code}): {create_response.text}") + pytest.fail( + f"Failed to create knowledge database (status={create_response.status_code}): {create_response.text}" + ) db_payload = create_response.json() db_id = db_payload["db_id"]