style: ruff format

This commit is contained in:
Wenjie Zhang 2025-12-14 22:44:18 +08:00
parent e1a0c6be2c
commit e662466fae
2 changed files with 26 additions and 37 deletions

View File

@ -14,11 +14,7 @@ evaluation = APIRouter(prefix="/evaluation", tags=["evaluation"])
@evaluation.get("/databases/{db_id}/benchmarks/{benchmark_id}") @evaluation.get("/databases/{db_id}/benchmarks/{benchmark_id}")
async def get_evaluation_benchmark_by_db( async def get_evaluation_benchmark_by_db(
db_id: str, db_id: str, benchmark_id: str, page: int = 1, page_size: int = 10, current_user: User = Depends(get_admin_user)
benchmark_id: str,
page: int = 1,
page_size: int = 10,
current_user: User = Depends(get_admin_user)
): ):
"""根据 db_id 获取评估基准详情(支持分页)""" """根据 db_id 获取评估基准详情(支持分页)"""
from src.services.evaluation_service import EvaluationService from src.services.evaluation_service import EvaluationService
@ -54,7 +50,6 @@ async def delete_evaluation_benchmark(benchmark_id: str, current_user: User = De
raise HTTPException(status_code=500, detail=f"删除评估基准失败: {str(e)}") raise HTTPException(status_code=500, detail=f"删除评估基准失败: {str(e)}")
@evaluation.get("/databases/{db_id}/results/{task_id}") @evaluation.get("/databases/{db_id}/results/{task_id}")
async def get_evaluation_results_by_db( async def get_evaluation_results_by_db(
db_id: str, db_id: str,
@ -62,7 +57,7 @@ async def get_evaluation_results_by_db(
page: int = 1, page: int = 1,
page_size: int = 20, page_size: int = 20,
error_only: bool = False, error_only: bool = False,
current_user: User = Depends(get_admin_user) current_user: User = Depends(get_admin_user),
): ):
"""获取评估结果(带 db_id支持分页""" """获取评估结果(带 db_id支持分页"""
from src.services.evaluation_service import EvaluationService from src.services.evaluation_service import EvaluationService

View File

@ -7,7 +7,6 @@ from datetime import datetime
from typing import Any from typing import Any
from server.services.tasker import TaskContext, tasker from server.services.tasker import TaskContext, tasker
from src import config
from src.knowledge import knowledge_base from src.knowledge import knowledge_base
from src.models import select_model from src.models import select_model
from src.utils import logger from src.utils import logger
@ -136,7 +135,9 @@ class EvaluationService:
logger.error(f"获取评估基准详情失败: {e}") logger.error(f"获取评估基准详情失败: {e}")
raise raise
async def get_benchmark_detail_by_db(self, db_id: str, benchmark_id: str, page: int = 1, page_size: int = 10) -> dict[str, Any]: async def get_benchmark_detail_by_db(
self, db_id: str, benchmark_id: str, page: int = 1, page_size: int = 10
) -> dict[str, Any]:
"""根据 db_id 获取评估基准详情(支持分页)""" """根据 db_id 获取评估基准详情(支持分页)"""
try: try:
kb_instance = knowledge_base.get_kb(db_id) kb_instance = knowledge_base.get_kb(db_id)
@ -174,7 +175,8 @@ class EvaluationService:
total_pages = (total_questions + page_size - 1) // page_size total_pages = (total_questions + page_size - 1) // page_size
meta_with_q = meta.copy() meta_with_q = meta.copy()
meta_with_q.update({ meta_with_q.update(
{
"questions": questions, "questions": questions,
"pagination": { "pagination": {
"current_page": page, "current_page": page,
@ -182,9 +184,10 @@ class EvaluationService:
"total_questions": total_questions, "total_questions": total_questions,
"total_pages": total_pages, "total_pages": total_pages,
"has_next": page < total_pages, "has_next": page < total_pages,
"has_prev": page > 1 "has_prev": page > 1,
},
} }
}) )
return meta_with_q return meta_with_q
except Exception as e: except Exception as e:
logger.error(f"获取评估基准详情失败: {e}") logger.error(f"获取评估基准详情失败: {e}")
@ -226,8 +229,8 @@ class EvaluationService:
return {"task_id": task_id, "message": "基准生成任务已提交"} return {"task_id": task_id, "message": "基准生成任务已提交"}
async def _generate_benchmark_task(self, context: TaskContext): async def _generate_benchmark_task(self, context: TaskContext):
import random
import math import math
import random
await context.set_progress(0, "初始化") await context.set_progress(0, "初始化")
@ -346,9 +349,8 @@ class EvaluationService:
prompt = ( prompt = (
"你将基于以下上下文生成一个可由上下文准确回答的问题与标准答案。" "你将基于以下上下文生成一个可由上下文准确回答的问题与标准答案。"
"仅返回一个JSON对象不要包含其他文字。" "仅返回一个JSON对象不要包含其他文字。"
"键为 query、gold_answer、gold_chunk_ids。gold_chunk_ids 必须是上述上下文片段的ID子集。\n\n上下文:\n" "键为 query、gold_answer、gold_chunk_ids。gold_chunk_ids 必须是上述上下文片段的ID子集。\n\n"
+ context_text "上下文:\n" + context_text + "\n"
+ "\n"
) )
try: try:
@ -356,6 +358,7 @@ class EvaluationService:
content = resp.content if resp else "" content = resp.content if resp else ""
import json_repair import json_repair
obj = json_repair.loads(content) obj = json_repair.loads(content)
q = obj.get("query") q = obj.get("query")
a = obj.get("gold_answer") a = obj.get("gold_answer")
@ -587,7 +590,7 @@ class EvaluationService:
prompt = ( prompt = (
f"基于以下上下文信息,请回答用户的问题。\n\n" f"基于以下上下文信息,请回答用户的问题。\n\n"
f"上下文信息:{context_text}\n\n" f"上下文信息:{context_text}\n\n"
f"用户问题:{question_data["query"]}\n\n" f"用户问题:{question_data['query']}\n\n"
"请根据上下文信息准确回答问题。\n\n" "请根据上下文信息准确回答问题。\n\n"
"如果上下文中缺少相关信息,请回答“信息不足,无法回答”。\n\n" "如果上下文中缺少相关信息,请回答“信息不足,无法回答”。\n\n"
) )
@ -762,12 +765,7 @@ class EvaluationService:
# 索引与回退逻辑已移除,统一通过 db_id 定位 # 索引与回退逻辑已移除,统一通过 db_id 定位
async def get_evaluation_results_by_db( async def get_evaluation_results_by_db(
self, self, db_id: str, task_id: str, page: int = 1, page_size: int = 20, error_only: bool = False
db_id: str,
task_id: str,
page: int = 1,
page_size: int = 20,
error_only: bool = False
) -> dict[str, Any]: ) -> dict[str, Any]:
result_file_path = os.path.join(self._get_result_dir(db_id), f"{task_id}.json") result_file_path = os.path.join(self._get_result_dir(db_id), f"{task_id}.json")
if not os.path.exists(result_file_path): if not os.path.exists(result_file_path):
@ -800,11 +798,7 @@ class EvaluationService:
# 检查检索指标是否明显偏低 # 检查检索指标是否明显偏低
metrics = item.get("metrics", {}) metrics = item.get("metrics", {})
has_low_recall = any( has_low_recall = any(metrics.get(k, 1.0) < 0.3 for k in metrics if k.startswith("recall@"))
metrics.get(k, 1.0) < 0.3
for k in metrics
if k.startswith("recall@")
)
if has_low_recall: if has_low_recall:
filtered_results.append(item) filtered_results.append(item)
all_results = filtered_results all_results = filtered_results
@ -831,8 +825,8 @@ class EvaluationService:
"page_size": page_size, "page_size": page_size,
"total": total, "total": total,
"total_pages": (total + page_size - 1) // page_size, "total_pages": (total + page_size - 1) // page_size,
"error_only": error_only "error_only": error_only,
} },
} }
# 非分页请求,返回完整数据(保持向后兼容) # 非分页请求,返回完整数据(保持向后兼容)