merge: pull origin/main and resolve batch size conflicts

This commit is contained in:
肖泽涛 2026-02-28 17:18:33 +08:00
commit 2b80e8c41f
3 changed files with 31 additions and 12 deletions

View File

@ -741,11 +741,11 @@ async def update_thread(
)
# ================================
# > === 附件管理分组 ===
# ================================
@chat.post("/thread/{thread_id}/attachments", response_model=AttachmentResponse)
async def upload_thread_attachment(
thread_id: str,

View File

@ -14,6 +14,7 @@ from src.agents.common.backends import create_agent_composite_backend
from src.agents.common.middlewares import RuntimeConfigMiddleware, SummaryOffloadMiddleware, save_attachments_to_fs
from src.agents.common.tools import get_tavily_search
from src.services.mcp_service import get_tools_from_all_servers
from src.utils import logger
from .context import DeepContext
@ -83,11 +84,14 @@ class DeepAgent(BaseAgent):
if tavily_search:
tools.append(tavily_search)
# Assert that search tool is available for DeepAgent
assert tools, (
"DeepAgent requires at least one search tool. "
"Please configure TAVILY_API_KEY environment variable to enable web search."
)
# # Assert that search tool is available for DeepAgent
# assert tools, (
# "DeepAgent requires at least one search tool. "
# "Please configure TAVILY_API_KEY environment variable to enable web search."
# )
if not tools:
logger.warning("No search tools configured, DeepAgent will work without web search")
tools = []
return tools
async def get_graph(self, **kwargs):

View File

@ -88,17 +88,32 @@ class BaseEmbeddingModel(ABC):
task_id = hashstr(messages)
self.embed_state[task_id] = {"status": "in-progress", "total": len(messages), "progress": 0}
tasks = []
# 保留原有逻辑:
# 使用 asyncio.gather 并发执行所有 embedding 批次请求:
# tasks = []
# for i in range(0, len(messages), batch_size):
# group_msg = messages[i : i + batch_size]
# tasks.append(self.aencode(group_msg))
# results = await asyncio.gather(*tasks)
# for res in results:
# data.extend(res)
# if task_id:
# self.embed_state[task_id]["progress"] = len(messages)
# self.embed_state[task_id]["status"] = "completed"
# return data
for i in range(0, len(messages), batch_size):
group_msg = messages[i : i + batch_size]
tasks.append(self.aencode(group_msg))
results = await asyncio.gather(*tasks)
for res in results:
logger.info(f"Async encoding [{i}/{len(messages)}] messages (bsz={batch_size})")
res = await self.aencode(group_msg)
data.extend(res)
if task_id:
self.embed_state[task_id]["progress"] = i + len(group_msg)
if task_id:
self.embed_state[task_id]["progress"] = len(messages)
self.embed_state[task_id]["status"] = "completed"
return data