feat: 增强对话线程的 agent_config_id 绑定逻辑,确保元数据一致性
This commit is contained in:
parent
3b5514e9b5
commit
5fe367b77b
@ -9,6 +9,6 @@ PROMPT = """
|
|||||||
- /home/gem/user-data/uploads/:用于存放用户上传的文件
|
- /home/gem/user-data/uploads/:用于存放用户上传的文件
|
||||||
|
|
||||||
如果启用了知识库,除了使用知识库工具之外,
|
如果启用了知识库,除了使用知识库工具之外,
|
||||||
当需要精准获取信息的时候,还可以直接访问知识库文件系统 (路径为 /home/gem/kbs/)来获取信息。
|
当需要精准获取信息的时候,或者 query_kb 中没有找到相关的内容,还可以直接访问知识库文件系统 (路径为 /home/gem/kbs/)来获取信息。
|
||||||
源文件可能无法解析,可以在 /home/gem/kbs/<db_name>/parsed/ 中找到解析后的 markdown 文件。
|
源文件可能无法解析,可以在 /home/gem/kbs/<db_name>/parsed/ 中找到解析后的 markdown 文件。
|
||||||
"""
|
"""
|
||||||
|
|||||||
@ -81,12 +81,27 @@ class ConversationRepository:
|
|||||||
metadata["attachments"] = list(metadata.get("attachments", []))
|
metadata["attachments"] = list(metadata.get("attachments", []))
|
||||||
return metadata
|
return metadata
|
||||||
|
|
||||||
|
def _normalize_agent_config_id(self, agent_config_id: int | None) -> int | None:
|
||||||
|
if agent_config_id is None:
|
||||||
|
return None
|
||||||
|
return int(agent_config_id)
|
||||||
|
|
||||||
async def _save_metadata(self, conversation: Conversation, metadata: dict) -> None:
|
async def _save_metadata(self, conversation: Conversation, metadata: dict) -> None:
|
||||||
conversation.extra_metadata = metadata
|
conversation.extra_metadata = metadata
|
||||||
conversation.updated_at = utc_now_naive()
|
conversation.updated_at = utc_now_naive()
|
||||||
await self.db.commit()
|
await self.db.commit()
|
||||||
await self.db.refresh(conversation)
|
await self.db.refresh(conversation)
|
||||||
|
|
||||||
|
async def bind_agent_config(self, thread_id: str, agent_config_id: int) -> Conversation | None:
|
||||||
|
conversation = await self.get_conversation_by_thread_id(thread_id)
|
||||||
|
if not conversation:
|
||||||
|
return None
|
||||||
|
|
||||||
|
metadata = self._ensure_metadata(conversation)
|
||||||
|
metadata["agent_config_id"] = self._normalize_agent_config_id(agent_config_id)
|
||||||
|
await self._save_metadata(conversation, metadata)
|
||||||
|
return conversation
|
||||||
|
|
||||||
async def add_message(
|
async def add_message(
|
||||||
self,
|
self,
|
||||||
conversation_id: int,
|
conversation_id: int,
|
||||||
|
|||||||
@ -79,6 +79,10 @@ async def create_agent_run_view(
|
|||||||
conversation = await conv_repo.get_conversation_by_thread_id(thread_id)
|
conversation = await conv_repo.get_conversation_by_thread_id(thread_id)
|
||||||
if not conversation or conversation.user_id != str(current_user_id) or conversation.status == "deleted":
|
if not conversation or conversation.user_id != str(current_user_id) or conversation.status == "deleted":
|
||||||
raise HTTPException(status_code=404, detail="对话线程不存在")
|
raise HTTPException(status_code=404, detail="对话线程不存在")
|
||||||
|
if (conversation.extra_metadata or {}).get("agent_config_id") != int(agent_config_id):
|
||||||
|
conversation = await conv_repo.bind_agent_config(thread_id, agent_config_id)
|
||||||
|
if not conversation:
|
||||||
|
raise HTTPException(status_code=404, detail="对话线程不存在")
|
||||||
|
|
||||||
request_id = str((meta or {}).get("request_id") or uuid.uuid4())
|
request_id = str((meta or {}).get("request_id") or uuid.uuid4())
|
||||||
config = {
|
config = {
|
||||||
|
|||||||
@ -371,6 +371,27 @@ async def check_and_handle_interrupts(
|
|||||||
logger.error(traceback.format_exc())
|
logger.error(traceback.format_exc())
|
||||||
|
|
||||||
|
|
||||||
|
async def _ensure_thread_bound_agent_config(
|
||||||
|
*,
|
||||||
|
conv_repo: ConversationRepository,
|
||||||
|
thread_id: str,
|
||||||
|
user_id: str,
|
||||||
|
agent_id: str,
|
||||||
|
agent_config_id: int,
|
||||||
|
) -> None:
|
||||||
|
conversation = await conv_repo.get_conversation_by_thread_id(thread_id)
|
||||||
|
if not conversation:
|
||||||
|
conversation = await conv_repo.create_conversation(
|
||||||
|
user_id=user_id,
|
||||||
|
agent_id=agent_id,
|
||||||
|
thread_id=thread_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
current_agent_config_id = (conversation.extra_metadata or {}).get("agent_config_id")
|
||||||
|
if current_agent_config_id != int(agent_config_id):
|
||||||
|
await conv_repo.bind_agent_config(thread_id, agent_config_id)
|
||||||
|
|
||||||
|
|
||||||
async def agent_chat(
|
async def agent_chat(
|
||||||
*,
|
*,
|
||||||
query: str,
|
query: str,
|
||||||
@ -462,6 +483,13 @@ async def agent_chat(
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
conv_repo = ConversationRepository(db)
|
conv_repo = ConversationRepository(db)
|
||||||
|
await _ensure_thread_bound_agent_config(
|
||||||
|
conv_repo=conv_repo,
|
||||||
|
thread_id=thread_id,
|
||||||
|
user_id=user_id,
|
||||||
|
agent_id=agent_id,
|
||||||
|
agent_config_id=agent_config_id,
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await conv_repo.add_message_by_thread_id(
|
await conv_repo.add_message_by_thread_id(
|
||||||
@ -631,6 +659,13 @@ async def stream_agent_chat(
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
conv_repo = ConversationRepository(db)
|
conv_repo = ConversationRepository(db)
|
||||||
|
await _ensure_thread_bound_agent_config(
|
||||||
|
conv_repo=conv_repo,
|
||||||
|
thread_id=thread_id,
|
||||||
|
user_id=user_id,
|
||||||
|
agent_id=agent_id,
|
||||||
|
agent_config_id=agent_config_id,
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await conv_repo.add_message_by_thread_id(
|
await conv_repo.add_message_by_thread_id(
|
||||||
|
|||||||
@ -287,6 +287,7 @@ async def create_thread_view(
|
|||||||
"title": conversation.title,
|
"title": conversation.title,
|
||||||
"created_at": conversation.created_at.isoformat(),
|
"created_at": conversation.created_at.isoformat(),
|
||||||
"updated_at": conversation.updated_at.isoformat(),
|
"updated_at": conversation.updated_at.isoformat(),
|
||||||
|
"metadata": conversation.extra_metadata or {},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@ -316,6 +317,7 @@ async def list_threads_view(
|
|||||||
"is_pinned": bool(conv.is_pinned),
|
"is_pinned": bool(conv.is_pinned),
|
||||||
"created_at": conv.created_at.isoformat(),
|
"created_at": conv.created_at.isoformat(),
|
||||||
"updated_at": conv.updated_at.isoformat(),
|
"updated_at": conv.updated_at.isoformat(),
|
||||||
|
"metadata": conv.extra_metadata or {},
|
||||||
}
|
}
|
||||||
for conv in conversations
|
for conv in conversations
|
||||||
]
|
]
|
||||||
@ -356,6 +358,7 @@ async def update_thread_view(
|
|||||||
"is_pinned": bool(updated_conv.is_pinned),
|
"is_pinned": bool(updated_conv.is_pinned),
|
||||||
"created_at": updated_conv.created_at.isoformat(),
|
"created_at": updated_conv.created_at.isoformat(),
|
||||||
"updated_at": updated_conv.updated_at.isoformat(),
|
"updated_at": updated_conv.updated_at.isoformat(),
|
||||||
|
"metadata": updated_conv.extra_metadata or {},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -231,6 +231,7 @@ class Conversation(Base):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def to_dict(self) -> dict[str, Any]:
|
def to_dict(self) -> dict[str, Any]:
|
||||||
|
metadata = self.extra_metadata or {}
|
||||||
return {
|
return {
|
||||||
"id": self.id,
|
"id": self.id,
|
||||||
"thread_id": self.thread_id,
|
"thread_id": self.thread_id,
|
||||||
@ -241,7 +242,7 @@ class Conversation(Base):
|
|||||||
"is_pinned": bool(self.is_pinned),
|
"is_pinned": bool(self.is_pinned),
|
||||||
"created_at": format_utc_datetime(self.created_at),
|
"created_at": format_utc_datetime(self.created_at),
|
||||||
"updated_at": format_utc_datetime(self.updated_at),
|
"updated_at": format_utc_datetime(self.updated_at),
|
||||||
"metadata": self.extra_metadata or {},
|
"metadata": metadata,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -636,6 +636,7 @@ class ThreadResponse(BaseModel):
|
|||||||
is_pinned: bool = False
|
is_pinned: bool = False
|
||||||
created_at: str
|
created_at: str
|
||||||
updated_at: str
|
updated_at: str
|
||||||
|
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
class AttachmentResponse(BaseModel):
|
class AttachmentResponse(BaseModel):
|
||||||
|
|||||||
@ -98,6 +98,16 @@ async def test_chat_agent_sync_with_thread_id(test_client, admin_headers):
|
|||||||
f"thread_id mismatch: expected {thread_id}, got {payload.get('thread_id')}"
|
f"thread_id mismatch: expected {thread_id}, got {payload.get('thread_id')}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
threads_response = await test_client.get("/api/chat/threads", headers=admin_headers)
|
||||||
|
assert threads_response.status_code == 200, threads_response.text
|
||||||
|
threads = threads_response.json()
|
||||||
|
target_thread = next((item for item in threads if item.get("id") == thread_id), None)
|
||||||
|
assert target_thread is not None, f"thread not found in thread list: {thread_id}"
|
||||||
|
assert (target_thread.get("metadata") or {}).get("agent_config_id") == agent_config_id, (
|
||||||
|
"agent_config_id mismatch: "
|
||||||
|
f"expected {agent_config_id}, got {(target_thread.get('metadata') or {}).get('agent_config_id')}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def test_chat_agent_sync_with_meta(test_client, admin_headers):
|
async def test_chat_agent_sync_with_meta(test_client, admin_headers):
|
||||||
"""测试非流式对话传递 meta 参数"""
|
"""测试非流式对话传递 meta 参数"""
|
||||||
|
|||||||
@ -20,3 +20,19 @@ def test_normalize_title_trims_spaces():
|
|||||||
normalized = repo._normalize_title(" hello world ")
|
normalized = repo._normalize_title(" hello world ")
|
||||||
|
|
||||||
assert normalized == "hello world"
|
assert normalized == "hello world"
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_agent_config_id_casts_to_int():
|
||||||
|
repo = ConversationRepository(None) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
normalized = repo._normalize_agent_config_id("12") # type: ignore[arg-type]
|
||||||
|
|
||||||
|
assert normalized == 12
|
||||||
|
|
||||||
|
|
||||||
|
def test_normalize_agent_config_id_allows_none():
|
||||||
|
repo = ConversationRepository(None) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
normalized = repo._normalize_agent_config_id(None)
|
||||||
|
|
||||||
|
assert normalized is None
|
||||||
|
|||||||
@ -626,12 +626,16 @@ const setThreadAgentConfigId = (threadId, agentConfigId) => {
|
|||||||
if (!threadId) return
|
if (!threadId) return
|
||||||
const thread = threads.value.find((item) => item.id === threadId)
|
const thread = threads.value.find((item) => item.id === threadId)
|
||||||
if (thread) {
|
if (thread) {
|
||||||
thread.agent_config_id = agentConfigId ?? null
|
thread.metadata = {
|
||||||
|
...(thread.metadata || {}),
|
||||||
|
agent_config_id: agentConfigId ?? null
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const syncSelectedConfigForThread = async (thread) => {
|
const syncSelectedConfigForThread = async (thread) => {
|
||||||
if (!thread?.agent_config_id) return
|
const threadAgentConfigId = thread?.metadata?.agent_config_id
|
||||||
|
if (!threadAgentConfigId) return
|
||||||
|
|
||||||
const targetAgentId = thread.agent_id || currentAgentId.value
|
const targetAgentId = thread.agent_id || currentAgentId.value
|
||||||
if (!targetAgentId) return
|
if (!targetAgentId) return
|
||||||
@ -641,8 +645,8 @@ const syncSelectedConfigForThread = async (thread) => {
|
|||||||
await agentStore.fetchAgentConfigs(targetAgentId)
|
await agentStore.fetchAgentConfigs(targetAgentId)
|
||||||
}
|
}
|
||||||
|
|
||||||
if (selectedAgentConfigId.value !== thread.agent_config_id) {
|
if (selectedAgentConfigId.value !== threadAgentConfigId) {
|
||||||
await agentStore.selectAgentConfig(thread.agent_config_id)
|
await agentStore.selectAgentConfig(threadAgentConfigId)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user