feat: 添加 CLAUDE.md 文件以提供项目开发指导
- 新增 CLAUDE.md 文件,包含项目概述、开发环境管理、项目架构、开发规范、特色功能模块及调试和故障排除指南。 - 详细描述了使用 Docker Compose 管理开发环境的步骤和常用命令。 - 介绍了项目的技术栈、核心目录结构及主要服务,确保开发者能够快速上手。
This commit is contained in:
parent
0e92aaa99a
commit
67ddbe6f40
158
CLAUDE.md
158
CLAUDE.md
@ -1 +1,159 @@
|
|||||||
|
# CLAUDE.md
|
||||||
|
|
||||||
|
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
|
||||||
|
|
||||||
|
## 项目概述
|
||||||
|
|
||||||
|
Yuxi-Know 是一个基于大模型的智能知识库与知识图谱智能体开发平台,融合了 RAG 技术与知识图谱技术,基于 LangGraph v1 + Vue.js + FastAPI + LightRAG 架构构建。项目完全通过 Docker Compose 进行管理,支持热重载开发。
|
||||||
|
|
||||||
|
## 开发环境管理
|
||||||
|
|
||||||
|
### 核心原则
|
||||||
|
- 所有开发调试都应在运行的 Docker 容器环境中进行
|
||||||
|
- 使用 `docker compose up -d` 启动完整开发环境
|
||||||
|
- api-dev 和 web-dev 服务均配置热重载,本地修改代码后无需重启容器
|
||||||
|
- 先检查 `docker ps` 确认服务状态,使用 `docker logs api-dev --tail 100` 查看日志
|
||||||
|
|
||||||
|
### 常用命令
|
||||||
|
|
||||||
|
#### 项目启动和停止
|
||||||
|
```bash
|
||||||
|
# 启动所有服务
|
||||||
|
make start
|
||||||
|
# 或
|
||||||
|
docker compose up -d
|
||||||
|
|
||||||
|
# 停止所有服务
|
||||||
|
make stop
|
||||||
|
# 或
|
||||||
|
docker compose down
|
||||||
|
|
||||||
|
# 查看日志
|
||||||
|
make logs
|
||||||
|
# 或
|
||||||
|
docker logs --tail=50 api-dev
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 后端开发(在 Docker 容器内)
|
||||||
|
```bash
|
||||||
|
# 代码检查和格式化
|
||||||
|
make lint # 检查代码规范
|
||||||
|
make format # 格式化代码
|
||||||
|
make format_diff # 查看格式化差异
|
||||||
|
|
||||||
|
# 运行测试
|
||||||
|
make router-tests # 运行 API 路由测试
|
||||||
|
|
||||||
|
# 直接在容器内执行命令
|
||||||
|
docker compose exec api uv run python your_script.py
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 前端开发
|
||||||
|
```bash
|
||||||
|
# 在 web 目录下执行
|
||||||
|
pnpm run dev # 开发模式(已通过 Docker 配置)
|
||||||
|
pnpm run build # 构建生产版本
|
||||||
|
pnpm run lint # ESLint 检查
|
||||||
|
pnpm run format # Prettier 格式化
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 文档开发
|
||||||
|
```bash
|
||||||
|
# 在 docs 目录下执行
|
||||||
|
pnpm run docs:dev # 开发文档
|
||||||
|
pnpm run docs:build # 构建文档
|
||||||
|
```
|
||||||
|
|
||||||
|
## 项目架构
|
||||||
|
|
||||||
|
### 技术栈
|
||||||
|
- **后端**: FastAPI + Python 3.11+,使用 uv 管理依赖
|
||||||
|
- **前端**: Vue 3.5 + Vite 7 + Ant Design Vue + Pinia
|
||||||
|
- **AI框架**: LangChain v1 + LangGraph v1 + LightRAG
|
||||||
|
- **数据库**: Neo4j (图数据库) + Milvus (向量数据库) + MinIO (对象存储)
|
||||||
|
- **容器化**: Docker Compose 多服务编排
|
||||||
|
|
||||||
|
### 核心目录结构
|
||||||
|
```
|
||||||
|
├── server/ # FastAPI 后端服务
|
||||||
|
├── web/ # Vue.js 前端应用
|
||||||
|
├── src/ # 核心业务逻辑代码
|
||||||
|
├── docs/ # 文档中心 (VitePress)
|
||||||
|
├── test/ # 测试代码
|
||||||
|
├── docker/ # Docker 配置文件
|
||||||
|
├── scripts/ # 脚本工具
|
||||||
|
├── saves/ # 数据保存目录
|
||||||
|
├── models/ # 模型文件目录
|
||||||
|
└── docker-compose.yml
|
||||||
|
```
|
||||||
|
|
||||||
|
### 主要服务
|
||||||
|
- `api-dev`: FastAPI 后端服务 (端口 5050)
|
||||||
|
- `web-dev`: Vue.js 前端服务 (端口 5173)
|
||||||
|
- `graph`: Neo4j 图数据库 (端口 7474, 7687)
|
||||||
|
- `milvus`: Milvus 向量数据库 (端口 19530)
|
||||||
|
- `minio`: MinIO 对象存储 (端口 9000, 9001)
|
||||||
|
- `mineru`: MinerU 文档解析服务 (可选)
|
||||||
|
- `paddlex`: PaddleX OCR 服务 (可选)
|
||||||
|
|
||||||
|
## 开发规范
|
||||||
|
|
||||||
|
### 前端开发规范
|
||||||
|
- **API 接口**: 所有 API 接口定义在 `web/src/apis` 下,继承自 `apiGet/apiPost/apiRequest`
|
||||||
|
- **图标**: 从 `@ant-design/icons-vue` 或 `lucide-vue-next` 选取
|
||||||
|
- **样式**: 使用 Less,优先采用 `web/src/assets/css/base.css` 中的颜色
|
||||||
|
- **UI风格**: 简洁一致,避免悬停位移、过度阴影和渐变色
|
||||||
|
- **组件库**: 基于 Ant Design Vue,保持组件一致性
|
||||||
|
|
||||||
|
### 后端开发规范
|
||||||
|
- **包管理**: 使用 uv 管理依赖,调试时使用 `uv run`
|
||||||
|
- **代码规范**: 符合 Pythonic 风格,支持 Python 3.12+ 语法
|
||||||
|
- **代码质量**: 使用 `make lint` 检查,`make format` 格式化
|
||||||
|
- **API设计**: 遵循 RESTful 规范,使用 Pydantic 进行数据验证
|
||||||
|
|
||||||
|
### 通用开发规范
|
||||||
|
- **文档更新**: 代码更新后同步更新 `docs/latest` 中的相关文档
|
||||||
|
- **测试**: 测试脚本放在 `test/` 目录,从 Docker 容器中运行
|
||||||
|
- **环境变量**: 配置通过 `.env` 文件管理
|
||||||
|
- **数据安全**: 敏感数据不提交到版本控制
|
||||||
|
|
||||||
|
## 特色功能模块
|
||||||
|
|
||||||
|
### 智能体系统
|
||||||
|
- 基于 LangGraph v1 的智能体框架
|
||||||
|
- 支持多代理协作和工具调用
|
||||||
|
- 提供完整的智能体开发套件
|
||||||
|
|
||||||
|
### 知识管理
|
||||||
|
- 多模态文档解析 (PDF、Word、图片等)
|
||||||
|
- 知识图谱自动构建和可视化
|
||||||
|
- 向量化存储和检索
|
||||||
|
|
||||||
|
### 数据处理
|
||||||
|
- MinerU 文档解析集成
|
||||||
|
- PaddleX OCR 文字识别
|
||||||
|
- 支持多种模型提供商 (OpenAI、DeepSeek、阿里云等)
|
||||||
|
|
||||||
|
## 调试和故障排除
|
||||||
|
|
||||||
|
### 常见问题排查
|
||||||
|
1. **服务启动失败**: 检查端口占用和 Docker 服务状态
|
||||||
|
2. **API 连接问题**: 确认 `VITE_API_URL` 环境变量配置
|
||||||
|
3. **模型加载问题**: 检查 `models/` 目录和权限设置
|
||||||
|
4. **数据库连接**: 确认 Neo4j 和 Milvus 服务健康状态
|
||||||
|
|
||||||
|
### 日志查看
|
||||||
|
```bash
|
||||||
|
# 查看各服务日志
|
||||||
|
docker logs api-dev --tail 100
|
||||||
|
docker logs web-dev --tail 50
|
||||||
|
docker logs graph --tail 50
|
||||||
|
docker logs milvus --tail 50
|
||||||
|
```
|
||||||
|
|
||||||
|
### 性能监控
|
||||||
|
- 使用健康检查端点 `/api/system/health`
|
||||||
|
- 监控 GPU 使用情况 (如果使用 GPU 服务)
|
||||||
|
- 检查内存和磁盘使用情况
|
||||||
|
|
||||||
You MUST read `./AGENTS.md`
|
You MUST read `./AGENTS.md`
|
||||||
@ -180,31 +180,41 @@ def _serialize_attachment(record: dict) -> dict:
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
async def save_partial_message(conv_mgr, thread_id, full_msg, error_type="interrupted"):
|
async def save_partial_message(conv_mgr, thread_id, full_msg=None, error_message=None, error_type="interrupted"):
|
||||||
"""
|
"""
|
||||||
保存部分生成的AI消息到数据库
|
统一保存AI消息到数据库的函数
|
||||||
当streaming被中断时,用于保存已生成的部分内容
|
|
||||||
"""
|
|
||||||
if not full_msg:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
Args:
|
||||||
|
conv_mgr: 对话管理器
|
||||||
|
thread_id: 线程ID
|
||||||
|
full_msg: 完整的AI消息对象(可选)
|
||||||
|
error_message: 纯错误消息文本(当full_msg为空时使用)
|
||||||
|
error_type: 错误类型标识
|
||||||
|
"""
|
||||||
try:
|
try:
|
||||||
msg_dict = full_msg.model_dump() if hasattr(full_msg, "model_dump") else {}
|
if full_msg:
|
||||||
content = full_msg.content if hasattr(full_msg, "content") else str(full_msg)
|
# 保存部分生成的AI消息
|
||||||
|
msg_dict = full_msg.model_dump() if hasattr(full_msg, "model_dump") else {}
|
||||||
|
content = full_msg.content if hasattr(full_msg, "content") else str(full_msg)
|
||||||
|
extra_metadata = msg_dict | {"error_type": error_type}
|
||||||
|
else:
|
||||||
|
# 保存纯错误消息
|
||||||
|
content = error_message or f"发生错误: {error_type}"
|
||||||
|
extra_metadata = {"error_type": error_type, "is_error": True}
|
||||||
|
|
||||||
saved_msg = conv_mgr.add_message_by_thread_id(
|
saved_msg = conv_mgr.add_message_by_thread_id(
|
||||||
thread_id=thread_id,
|
thread_id=thread_id,
|
||||||
role="assistant",
|
role="assistant",
|
||||||
content=content,
|
content=content,
|
||||||
message_type="text",
|
message_type="text",
|
||||||
extra_metadata=msg_dict | {"error_type": error_type},
|
extra_metadata=extra_metadata,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.info(f"Saved partial message due to {error_type}: {saved_msg.id}")
|
logger.info(f"Saved message due to {error_type}: {saved_msg.id}")
|
||||||
return saved_msg
|
return saved_msg
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error saving partial message: {e}")
|
logger.error(f"Error saving message: {e}")
|
||||||
logger.error(traceback.format_exc())
|
logger.error(traceback.format_exc())
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@ -367,7 +377,7 @@ async def get_agent(current_user: User = Depends(get_required_user)):
|
|||||||
"examples": agent_info.get("examples", []),
|
"examples": agent_info.get("examples", []),
|
||||||
"configurable_items": agent_info.get("configurable_items", []),
|
"configurable_items": agent_info.get("configurable_items", []),
|
||||||
"has_checkpointer": agent_info.get("has_checkpointer", False),
|
"has_checkpointer": agent_info.get("has_checkpointer", False),
|
||||||
"capabilities": agent_info.get("capabilities", []) # 智能体能力列表
|
"capabilities": agent_info.get("capabilities", []), # 智能体能力列表
|
||||||
}
|
}
|
||||||
for agent_info in agents_info
|
for agent_info in agents_info
|
||||||
]
|
]
|
||||||
@ -513,16 +523,19 @@ async def chat_agent(
|
|||||||
# 客户端主动中断连接,检查中断并保存已生成的部分内容
|
# 客户端主动中断连接,检查中断并保存已生成的部分内容
|
||||||
logger.warning(f"Client disconnected, cancelling stream: {e}")
|
logger.warning(f"Client disconnected, cancelling stream: {e}")
|
||||||
|
|
||||||
|
# 保存中断消息到数据库
|
||||||
# 如果有手动维护的 full_msg,直接保存到数据库
|
new_db = db_manager.get_session()
|
||||||
if full_msg:
|
try:
|
||||||
# 创建新的 db session,因为原 session 可能已关闭
|
new_conv_manager = ConversationManager(new_db)
|
||||||
new_db = db_manager.get_session()
|
await save_partial_message(
|
||||||
try:
|
new_conv_manager,
|
||||||
new_conv_manager = ConversationManager(new_db)
|
thread_id,
|
||||||
await save_partial_message(new_conv_manager, thread_id, full_msg, "interrupted")
|
full_msg=full_msg,
|
||||||
finally:
|
error_message="对话已中断" if not full_msg else None,
|
||||||
new_db.close()
|
error_type="interrupted",
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
new_db.close()
|
||||||
|
|
||||||
# 通知前端中断(可能发送不到,但用于一致性)
|
# 通知前端中断(可能发送不到,但用于一致性)
|
||||||
yield make_chunk(status="interrupted", message="对话已中断", meta=meta)
|
yield make_chunk(status="interrupted", message="对话已中断", meta=meta)
|
||||||
@ -530,16 +543,19 @@ async def chat_agent(
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error streaming messages: {e}, {traceback.format_exc()}")
|
logger.error(f"Error streaming messages: {e}, {traceback.format_exc()}")
|
||||||
|
|
||||||
|
# 保存错误消息到数据库
|
||||||
# 如果有手动维护的 full_msg,直接保存到数据库
|
new_db = db_manager.get_session()
|
||||||
if full_msg:
|
try:
|
||||||
# 创建新的 db session,因为原 session 可能已关闭
|
new_conv_manager = ConversationManager(new_db)
|
||||||
new_db = db_manager.get_session()
|
await save_partial_message(
|
||||||
try:
|
new_conv_manager,
|
||||||
new_conv_manager = ConversationManager(new_db)
|
thread_id,
|
||||||
await save_partial_message(new_conv_manager, thread_id, full_msg, "unexpected_error")
|
full_msg=full_msg,
|
||||||
finally:
|
error_message=f"Error streaming messages: {e}" if not full_msg else None,
|
||||||
new_db.close()
|
error_type="unexpected_error",
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
new_db.close()
|
||||||
|
|
||||||
yield make_chunk(message=f"Error streaming messages: {e}", status="error")
|
yield make_chunk(message=f"Error streaming messages: {e}", status="error")
|
||||||
|
|
||||||
@ -644,28 +660,61 @@ async def resume_agent_chat(
|
|||||||
resume_command, context=context, config={"configurable": input_context}, stream_mode="messages"
|
resume_command, context=context, config={"configurable": input_context}, stream_mode="messages"
|
||||||
)
|
)
|
||||||
|
|
||||||
async for msg, metadata in stream_source:
|
try:
|
||||||
# 确保msg有正确的ID结构
|
async for msg, metadata in stream_source:
|
||||||
msg_dict = msg.model_dump()
|
# 确保msg有正确的ID结构
|
||||||
if "id" not in msg_dict:
|
msg_dict = msg.model_dump()
|
||||||
msg_dict["id"] = str(uuid.uuid4())
|
if "id" not in msg_dict:
|
||||||
|
msg_dict["id"] = str(uuid.uuid4())
|
||||||
|
|
||||||
yield make_resume_chunk(
|
yield make_resume_chunk(
|
||||||
content=getattr(msg, "content", ""), msg=msg_dict, metadata=metadata, status="loading"
|
content=getattr(msg, "content", ""), msg=msg_dict, metadata=metadata, status="loading"
|
||||||
|
)
|
||||||
|
|
||||||
|
meta["time_cost"] = asyncio.get_event_loop().time() - start_time
|
||||||
|
yield make_resume_chunk(status="finished", meta=meta)
|
||||||
|
|
||||||
|
# 保存消息到数据库
|
||||||
|
langgraph_config = {"configurable": input_context}
|
||||||
|
conv_manager = ConversationManager(db)
|
||||||
|
await save_messages_from_langgraph_state(
|
||||||
|
agent_instance=agent,
|
||||||
|
thread_id=thread_id,
|
||||||
|
conv_mgr=conv_manager,
|
||||||
|
config_dict=langgraph_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
meta["time_cost"] = asyncio.get_event_loop().time() - start_time
|
except (asyncio.CancelledError, ConnectionError) as e:
|
||||||
yield make_resume_chunk(status="finished", meta=meta)
|
# 客户端主动中断连接
|
||||||
|
logger.warning(f"Client disconnected during resume: {e}")
|
||||||
|
|
||||||
# 保存消息到数据库
|
# 保存中断消息到数据库
|
||||||
langgraph_config = {"configurable": input_context}
|
new_db = db_manager.get_session()
|
||||||
conv_manager = ConversationManager(db)
|
try:
|
||||||
await save_messages_from_langgraph_state(
|
new_conv_manager = ConversationManager(new_db)
|
||||||
agent_instance=agent,
|
await save_partial_message(
|
||||||
thread_id=thread_id,
|
new_conv_manager, thread_id, error_message="对话恢复已中断", error_type="resume_interrupted"
|
||||||
conv_mgr=conv_manager,
|
)
|
||||||
config_dict=langgraph_config,
|
finally:
|
||||||
)
|
new_db.close()
|
||||||
|
|
||||||
|
yield make_resume_chunk(status="interrupted", message="对话恢复已中断", meta=meta)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
# 处理其他异常
|
||||||
|
logger.error(f"Error during resume: {e}, {traceback.format_exc()}")
|
||||||
|
|
||||||
|
# 保存错误消息到数据库
|
||||||
|
new_db = db_manager.get_session()
|
||||||
|
try:
|
||||||
|
new_conv_manager = ConversationManager(new_db)
|
||||||
|
await save_partial_message(
|
||||||
|
new_conv_manager, thread_id, error_message=f"Error during resume: {e}", error_type="resume_error"
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
new_db.close()
|
||||||
|
|
||||||
|
yield make_resume_chunk(message=f"Error during resume: {e}", status="error")
|
||||||
|
|
||||||
return StreamingResponse(stream_resume(), media_type="application/json")
|
return StreamingResponse(stream_resume(), media_type="application/json")
|
||||||
|
|
||||||
|
|||||||
@ -247,12 +247,14 @@ async def add_documents(
|
|||||||
error_type = "timeout" if isinstance(doc_error, TimeoutError) else "processing_error"
|
error_type = "timeout" if isinstance(doc_error, TimeoutError) else "processing_error"
|
||||||
error_msg = "处理超时" if isinstance(doc_error, TimeoutError) else "处理失败"
|
error_msg = "处理超时" if isinstance(doc_error, TimeoutError) else "处理失败"
|
||||||
|
|
||||||
processed_items.append({
|
processed_items.append(
|
||||||
"item": item,
|
{
|
||||||
"status": "failed",
|
"item": item,
|
||||||
"error": f"{error_msg}: {str(doc_error)}",
|
"status": "failed",
|
||||||
"error_type": error_type
|
"error": f"{error_msg}: {str(doc_error)}",
|
||||||
})
|
"error_type": error_type,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
await context.set_progress(100.0, "任务已取消")
|
await context.set_progress(100.0, "任务已取消")
|
||||||
@ -262,13 +264,15 @@ async def add_documents(
|
|||||||
logger.exception(f"Task processing failed: {task_error}")
|
logger.exception(f"Task processing failed: {task_error}")
|
||||||
await context.set_progress(100.0, f"任务处理失败: {str(task_error)}")
|
await context.set_progress(100.0, f"任务处理失败: {str(task_error)}")
|
||||||
# 将所有未处理的文档标记为失败
|
# 将所有未处理的文档标记为失败
|
||||||
for item in items[len(processed_items):]:
|
for item in items[len(processed_items) :]:
|
||||||
processed_items.append({
|
processed_items.append(
|
||||||
"item": item,
|
{
|
||||||
"status": "failed",
|
"item": item,
|
||||||
"error": f"任务失败: {str(task_error)}",
|
"status": "failed",
|
||||||
"error_type": "task_failed"
|
"error": f"任务失败: {str(task_error)}",
|
||||||
})
|
"error_type": "task_failed",
|
||||||
|
}
|
||||||
|
)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
item_type = "URL" if content_type == "url" else "文件"
|
item_type = "URL" if content_type == "url" else "文件"
|
||||||
@ -600,8 +604,7 @@ async def get_knowledge_base_query_params(db_id: str, current_user: User = Depen
|
|||||||
"type": "select",
|
"type": "select",
|
||||||
"default": reranker_config.get("model", ""),
|
"default": reranker_config.get("model", ""),
|
||||||
"options": [
|
"options": [
|
||||||
{"label": info.name, "value": model_id}
|
{"label": info.name, "value": model_id} for model_id, info in config.reranker_names.items()
|
||||||
for model_id, info in config.reranker_names.items()
|
|
||||||
],
|
],
|
||||||
"description": "覆盖默认配置,选择用于本次查询的重排序模型",
|
"description": "覆盖默认配置,选择用于本次查询的重排序模型",
|
||||||
}
|
}
|
||||||
@ -675,8 +678,7 @@ async def get_knowledge_base_query_params(db_id: str, current_user: User = Depen
|
|||||||
"type": "select",
|
"type": "select",
|
||||||
"default": reranker_config.get("model", ""),
|
"default": reranker_config.get("model", ""),
|
||||||
"options": [
|
"options": [
|
||||||
{"label": info.name, "value": model_id}
|
{"label": info.name, "value": model_id} for model_id, info in config.reranker_names.items()
|
||||||
for model_id, info in config.reranker_names.items()
|
|
||||||
],
|
],
|
||||||
"description": "覆盖默认配置,选择用于本次查询的重排序模型",
|
"description": "覆盖默认配置,选择用于本次查询的重排序模型",
|
||||||
}
|
}
|
||||||
|
|||||||
@ -51,8 +51,6 @@ async def update_config_batch(items: dict = Body(...), current_user: User = Depe
|
|||||||
return config.dump_config()
|
return config.dump_config()
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@system.get("/logs")
|
@system.get("/logs")
|
||||||
def get_system_logs(current_user: User = Depends(get_admin_user)):
|
def get_system_logs(current_user: User = Depends(get_admin_user)):
|
||||||
"""获取系统日志"""
|
"""获取系统日志"""
|
||||||
|
|||||||
@ -2,6 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import importlib.util
|
import importlib.util
|
||||||
import os
|
import os
|
||||||
|
import tomllib as tomli
|
||||||
from abc import abstractmethod
|
from abc import abstractmethod
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
@ -9,8 +10,6 @@ from langgraph.checkpoint.memory import InMemorySaver
|
|||||||
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite
|
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver, aiosqlite
|
||||||
from langgraph.graph.state import CompiledStateGraph
|
from langgraph.graph.state import CompiledStateGraph
|
||||||
|
|
||||||
import tomllib as tomli
|
|
||||||
|
|
||||||
from src import config as sys_config
|
from src import config as sys_config
|
||||||
from src.agents.common.context import BaseContext
|
from src.agents.common.context import BaseContext
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
@ -163,7 +162,6 @@ class BaseAgent:
|
|||||||
"""获取异步存储实例"""
|
"""获取异步存储实例"""
|
||||||
return AsyncSqliteSaver(await self.get_async_conn())
|
return AsyncSqliteSaver(await self.get_async_conn())
|
||||||
|
|
||||||
|
|
||||||
def load_metadata(self) -> dict:
|
def load_metadata(self) -> dict:
|
||||||
"""Load metadata from metadata.toml file in the agent's source directory."""
|
"""Load metadata from metadata.toml file in the agent's source directory."""
|
||||||
if self._metadata_cache is not None:
|
if self._metadata_cache is not None:
|
||||||
|
|||||||
@ -1,6 +1,6 @@
|
|||||||
|
from .attachment_middleware import inject_attachment_context
|
||||||
from .context_middlewares import context_aware_prompt, context_based_model
|
from .context_middlewares import context_aware_prompt, context_based_model
|
||||||
from .dynamic_tool_middleware import DynamicToolMiddleware
|
from .dynamic_tool_middleware import DynamicToolMiddleware
|
||||||
from .attachment_middleware import inject_attachment_context
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"DynamicToolMiddleware",
|
"DynamicToolMiddleware",
|
||||||
|
|||||||
@ -40,8 +40,7 @@ def _build_attachment_prompt(attachments: Sequence[dict]) -> str | None:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
instructions = (
|
instructions = (
|
||||||
"以下为用户提供的附件内容,请综合这些文件与用户的新问题进行回答。"
|
"以下为用户提供的附件内容,请综合这些文件与用户的新问题进行回答。如附件与问题无关,可忽略附件内容:\n\n"
|
||||||
"如附件与问题无关,可忽略附件内容:\n\n"
|
|
||||||
)
|
)
|
||||||
return instructions + "\n\n".join(chunks)
|
return instructions + "\n\n".join(chunks)
|
||||||
|
|
||||||
|
|||||||
@ -14,9 +14,7 @@ def context_aware_prompt(request: ModelRequest) -> str:
|
|||||||
|
|
||||||
|
|
||||||
@wrap_model_call
|
@wrap_model_call
|
||||||
async def context_based_model(
|
async def context_based_model(request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse]) -> ModelResponse:
|
||||||
request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse]
|
|
||||||
) -> ModelResponse:
|
|
||||||
"""从 runtime context 动态选择模型"""
|
"""从 runtime context 动态选择模型"""
|
||||||
model_spec = request.runtime.context.model
|
model_spec = request.runtime.context.model
|
||||||
model = load_chat_model(model_spec)
|
model = load_chat_model(model_spec)
|
||||||
|
|||||||
@ -11,6 +11,7 @@ calc_agent = create_agent(
|
|||||||
system_prompt="你可以使用计算器工具,处理各种数学计算任务。",
|
system_prompt="你可以使用计算器工具,处理各种数学计算任务。",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@tool(name_or_callable="calc_agent_tool", description="进行计算任务,输入是数学表达式或描述,输出计算结果。")
|
@tool(name_or_callable="calc_agent_tool", description="进行计算任务,输入是数学表达式或描述,输出计算结果。")
|
||||||
async def calc_agent_tool(description: str) -> str:
|
async def calc_agent_tool(description: str) -> str:
|
||||||
"""CalcAgent 工具 - 使用子智能体 CalcAgent 进行计算任务"""
|
"""CalcAgent 工具 - 使用子智能体 CalcAgent 进行计算任务"""
|
||||||
|
|||||||
@ -88,7 +88,7 @@ def mysql_list_tables() -> str:
|
|||||||
|
|
||||||
all_table_names = "\n".join(table_names)
|
all_table_names = "\n".join(table_names)
|
||||||
result = f"数据库中的表:\n{all_table_names}"
|
result = f"数据库中的表:\n{all_table_names}"
|
||||||
if db_note := conn_manager.config.get('description'):
|
if db_note := conn_manager.config.get("description"):
|
||||||
result = f"数据库说明: {db_note}\n\n" + result
|
result = f"数据库说明: {db_note}\n\n" + result
|
||||||
logger.info(f"Retrieved {len(table_names)} tables from database")
|
logger.info(f"Retrieved {len(table_names)} tables from database")
|
||||||
return result
|
return result
|
||||||
|
|||||||
@ -11,10 +11,10 @@ from pydantic import BaseModel, Field
|
|||||||
from src import config, graph_base, knowledge_base
|
from src import config, graph_base, knowledge_base
|
||||||
from src.utils import logger
|
from src.utils import logger
|
||||||
|
|
||||||
|
|
||||||
search = TavilySearch(max_results=10)
|
search = TavilySearch(max_results=10)
|
||||||
search.metadata = {"name": "Tavily 网页搜索"}
|
search.metadata = {"name": "Tavily 网页搜索"}
|
||||||
|
|
||||||
|
|
||||||
@tool(name_or_callable="计算器", description="可以对给定的2个数字选择进行 add, subtract, multiply, divide 运算")
|
@tool(name_or_callable="计算器", description="可以对给定的2个数字选择进行 add, subtract, multiply, divide 运算")
|
||||||
def calculator(a: float, b: float, operation: str) -> float:
|
def calculator(a: float, b: float, operation: str) -> float:
|
||||||
try:
|
try:
|
||||||
@ -38,7 +38,7 @@ def calculator(a: float, b: float, operation: str) -> float:
|
|||||||
@tool(name_or_callable="人工审批工具(Debug)", description="请求人工审批工具,用于在执行重要操作前获得人类确认。")
|
@tool(name_or_callable="人工审批工具(Debug)", description="请求人工审批工具,用于在执行重要操作前获得人类确认。")
|
||||||
def get_approved_user_goal(
|
def get_approved_user_goal(
|
||||||
operation_description: str,
|
operation_description: str,
|
||||||
)->dict:
|
) -> dict:
|
||||||
"""
|
"""
|
||||||
请求人工审批,在执行重要操作前获得人类确认。
|
请求人工审批,在执行重要操作前获得人类确认。
|
||||||
|
|
||||||
@ -78,6 +78,7 @@ KG_QUERY_DESCRIPTION = """
|
|||||||
关键词(query),使用可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。
|
关键词(query),使用可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
@tool(name_or_callable="查询知识图谱", description=KG_QUERY_DESCRIPTION)
|
@tool(name_or_callable="查询知识图谱", description=KG_QUERY_DESCRIPTION)
|
||||||
def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge graph."]) -> Any:
|
def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge graph."]) -> Any:
|
||||||
"""使用这个工具可以查询知识图谱中包含的三元组信息。关键词(query),使用可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。"""
|
"""使用这个工具可以查询知识图谱中包含的三元组信息。关键词(query),使用可能帮助回答这个问题的关键词进行查询,不要直接使用用户的原始输入去查询。"""
|
||||||
|
|||||||
@ -4,7 +4,6 @@ from dataclasses import field
|
|||||||
|
|
||||||
from src.agents.common.context import BaseContext
|
from src.agents.common.context import BaseContext
|
||||||
|
|
||||||
|
|
||||||
DEEP_PROMPT = """你是一个能够处理复杂、多步骤任务的深度分析代理。
|
DEEP_PROMPT = """你是一个能够处理复杂、多步骤任务的深度分析代理。
|
||||||
|
|
||||||
你拥有以下资源:
|
你拥有以下资源:
|
||||||
|
|||||||
@ -1,17 +1,20 @@
|
|||||||
"""Deep Agent - 基于create_deep_agent的深度分析智能体"""
|
"""Deep Agent - 基于create_deep_agent的深度分析智能体"""
|
||||||
|
|
||||||
import os
|
import os
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
|
|
||||||
from deepagents import create_deep_agent
|
from deepagents import create_deep_agent
|
||||||
|
from tavily import TavilyClient
|
||||||
|
|
||||||
from src.agents.common import BaseAgent, load_chat_model
|
from src.agents.common import BaseAgent, load_chat_model
|
||||||
from src.agents.common.tools import search
|
|
||||||
from src.agents.common.middlewares import (
|
from src.agents.common.middlewares import (
|
||||||
context_aware_prompt,
|
context_aware_prompt,
|
||||||
context_based_model,
|
context_based_model,
|
||||||
inject_attachment_context,
|
inject_attachment_context,
|
||||||
)
|
)
|
||||||
|
from src.agents.common.tools import search
|
||||||
|
|
||||||
from .context import DeepContext
|
from .context import DeepContext
|
||||||
from tavily import TavilyClient
|
|
||||||
|
|
||||||
# 最佳实践是初始化客户端一次并复用它。
|
# 最佳实践是初始化客户端一次并复用它。
|
||||||
tavily_client = TavilyClient(api_key=os.environ["TAVILY_API_KEY"])
|
tavily_client = TavilyClient(api_key=os.environ["TAVILY_API_KEY"])
|
||||||
@ -169,8 +172,8 @@ research_instructions = """你是一位专家级研究员。你的工作是进
|
|||||||
使用此工具对给定的查询进行网络搜索。你可以指定结果数量、主题以及是否包含原始内容。
|
使用此工具对给定的查询进行网络搜索。你可以指定结果数量、主题以及是否包含原始内容。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
class DeepAgent(BaseAgent):
|
|
||||||
|
|
||||||
|
class DeepAgent(BaseAgent):
|
||||||
name = "深度分析智能体"
|
name = "深度分析智能体"
|
||||||
description = "具备规划、深度分析和子智能体协作能力的智能体,可以处理复杂的多步骤任务"
|
description = "具备规划、深度分析和子智能体协作能力的智能体,可以处理复杂的多步骤任务"
|
||||||
context_schema = DeepContext
|
context_schema = DeepContext
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
|
|
||||||
from langchain.agents import create_agent
|
from langchain.agents import create_agent
|
||||||
|
|
||||||
from src import config
|
from src import config
|
||||||
|
|||||||
@ -323,9 +323,7 @@ class ChromaKB(KnowledgeBase):
|
|||||||
try:
|
try:
|
||||||
rerank_start = time.time()
|
rerank_start = time.time()
|
||||||
documents_text = [chunk["content"] for chunk in retrieved_chunks]
|
documents_text = [chunk["content"] for chunk in retrieved_chunks]
|
||||||
rerank_scores = await reranker.acompute_score(
|
rerank_scores = await reranker.acompute_score([query_text, documents_text], normalize=True)
|
||||||
[query_text, documents_text], normalize=True
|
|
||||||
)
|
|
||||||
|
|
||||||
for chunk, rerank_score in zip(retrieved_chunks, rerank_scores):
|
for chunk, rerank_score in zip(retrieved_chunks, rerank_scores):
|
||||||
chunk["rerank_score"] = float(rerank_score)
|
chunk["rerank_score"] = float(rerank_score)
|
||||||
|
|||||||
@ -369,9 +369,7 @@ class MilvusKB(KnowledgeBase):
|
|||||||
try:
|
try:
|
||||||
rerank_start = time.time()
|
rerank_start = time.time()
|
||||||
documents_text = [chunk["content"] for chunk in retrieved_chunks]
|
documents_text = [chunk["content"] for chunk in retrieved_chunks]
|
||||||
rerank_scores = await reranker.acompute_score(
|
rerank_scores = await reranker.acompute_score([query_text, documents_text], normalize=True)
|
||||||
[query_text, documents_text], normalize=True
|
|
||||||
)
|
|
||||||
|
|
||||||
for chunk, rerank_score in zip(retrieved_chunks, rerank_scores):
|
for chunk, rerank_score in zip(retrieved_chunks, rerank_scores):
|
||||||
chunk["rerank_score"] = float(rerank_score)
|
chunk["rerank_score"] = float(rerank_score)
|
||||||
|
|||||||
@ -85,9 +85,7 @@ async def test_admin_can_create_vector_db_with_reranker(test_client, admin_heade
|
|||||||
assert reranker_config.get("enabled") is True
|
assert reranker_config.get("enabled") is True
|
||||||
assert reranker_config.get("model") == "siliconflow/BAAI/bge-reranker-v2-m3"
|
assert reranker_config.get("model") == "siliconflow/BAAI/bge-reranker-v2-m3"
|
||||||
|
|
||||||
params_response = await test_client.get(
|
params_response = await test_client.get(f"/api/knowledge/databases/{db_id}/query-params", headers=admin_headers)
|
||||||
f"/api/knowledge/databases/{db_id}/query-params", headers=admin_headers
|
|
||||||
)
|
|
||||||
assert params_response.status_code == 200, params_response.text
|
assert params_response.status_code == 200, params_response.text
|
||||||
|
|
||||||
params_payload = params_response.json()
|
params_payload = params_response.json()
|
||||||
|
|||||||
@ -13,9 +13,7 @@ async def test_reranker_list_requires_admin(test_client, standard_user):
|
|||||||
public_response = await test_client.get("/api/settings/rerankers")
|
public_response = await test_client.get("/api/settings/rerankers")
|
||||||
assert public_response.status_code == 401
|
assert public_response.status_code == 401
|
||||||
|
|
||||||
forbidden_response = await test_client.get(
|
forbidden_response = await test_client.get("/api/settings/rerankers", headers=standard_user["headers"])
|
||||||
"/api/settings/rerankers", headers=standard_user["headers"]
|
|
||||||
)
|
|
||||||
assert forbidden_response.status_code == 403
|
assert forbidden_response.status_code == 403
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -436,6 +436,7 @@ const cleanupThreadState = (threadId) => {
|
|||||||
|
|
||||||
// ==================== STREAM HANDLING LOGIC ====================
|
// ==================== STREAM HANDLING LOGIC ====================
|
||||||
const resetOnGoingConv = (threadId = null, preserveMessages = false) => {
|
const resetOnGoingConv = (threadId = null, preserveMessages = false) => {
|
||||||
|
console.log('🔄 [RESET] Resetting on going conversation:', threadId, preserveMessages);
|
||||||
if (threadId) {
|
if (threadId) {
|
||||||
// 清理指定线程的状态
|
// 清理指定线程的状态
|
||||||
const threadState = getThreadState(threadId);
|
const threadState = getThreadState(threadId);
|
||||||
@ -449,8 +450,8 @@ const resetOnGoingConv = (threadId = null, preserveMessages = false) => {
|
|||||||
// 延迟清空消息,给历史记录加载足够时间
|
// 延迟清空消息,给历史记录加载足够时间
|
||||||
setTimeout(() => {
|
setTimeout(() => {
|
||||||
if (threadState.onGoingConv) {
|
if (threadState.onGoingConv) {
|
||||||
threadState.onGoingConv = createOnGoingConvState();
|
threadState.onGoingConv = createOnGoingConvState();
|
||||||
}
|
}
|
||||||
}, 100);
|
}, 100);
|
||||||
} else {
|
} else {
|
||||||
threadState.onGoingConv = createOnGoingConvState();
|
threadState.onGoingConv = createOnGoingConvState();
|
||||||
@ -469,8 +470,8 @@ const resetOnGoingConv = (threadId = null, preserveMessages = false) => {
|
|||||||
if (preserveMessages) {
|
if (preserveMessages) {
|
||||||
setTimeout(() => {
|
setTimeout(() => {
|
||||||
if (threadState.onGoingConv) {
|
if (threadState.onGoingConv) {
|
||||||
threadState.onGoingConv = createOnGoingConvState();
|
threadState.onGoingConv = createOnGoingConvState();
|
||||||
}
|
}
|
||||||
}, 100);
|
}, 100);
|
||||||
} else {
|
} else {
|
||||||
threadState.onGoingConv = createOnGoingConvState();
|
threadState.onGoingConv = createOnGoingConvState();
|
||||||
@ -547,8 +548,8 @@ const _processStreamChunk = (chunk, threadId) => {
|
|||||||
}
|
}
|
||||||
fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId })
|
fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId })
|
||||||
.finally(() => {
|
.finally(() => {
|
||||||
resetOnGoingConv(threadId, true);
|
resetOnGoingConv(threadId, true);
|
||||||
});
|
});
|
||||||
return true;
|
return true;
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -646,11 +647,17 @@ const updateThread = async (threadId, title) => {
|
|||||||
};
|
};
|
||||||
|
|
||||||
// 获取线程消息
|
// 获取线程消息
|
||||||
const fetchThreadMessages = async ({ agentId, threadId }) => {
|
const fetchThreadMessages = async ({ agentId, threadId, delay = 0 }) => {
|
||||||
if (!threadId || !agentId) return;
|
if (!threadId || !agentId) return;
|
||||||
|
|
||||||
|
// 如果指定了延迟,等待指定时间(用于确保后端数据库事务提交)
|
||||||
|
if (delay > 0) {
|
||||||
|
await new Promise(resolve => setTimeout(resolve, delay));
|
||||||
|
}
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const response = await agentApi.getAgentHistory(agentId, threadId);
|
const response = await agentApi.getAgentHistory(agentId, threadId);
|
||||||
|
console.log('🔄 [FETCH] Thread messages:', response);
|
||||||
threadMessages.value[threadId] = response.history || [];
|
threadMessages.value[threadId] = response.history || [];
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
handleChatError(error, 'load');
|
handleChatError(error, 'load');
|
||||||
@ -963,7 +970,7 @@ const handleSendOrStop = async () => {
|
|||||||
|
|
||||||
// 中断后刷新消息历史,确保显示最新的状态
|
// 中断后刷新消息历史,确保显示最新的状态
|
||||||
try {
|
try {
|
||||||
await fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId });
|
await fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId, delay: 100 });
|
||||||
message.info('已中断对话生成');
|
message.info('已中断对话生成');
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('刷新消息历史失败:', error);
|
console.error('刷新消息历史失败:', error);
|
||||||
|
|||||||
@ -118,8 +118,11 @@
|
|||||||
<div v-if="providerConfig.allModels.length === 0" class="modal-no-models">
|
<div v-if="providerConfig.allModels.length === 0" class="modal-no-models">
|
||||||
<a-alert v-if="!modelStatus[providerConfig.provider]" type="warning" message="请在 .env 中配置对应的 APIKEY,并重新启动服务" />
|
<a-alert v-if="!modelStatus[providerConfig.provider]" type="warning" message="请在 .env 中配置对应的 APIKEY,并重新启动服务" />
|
||||||
<div v-else>
|
<div v-else>
|
||||||
<a-alert type="warning" message="该提供商暂未适配获取模型列表的方法,如需添加模型,请编辑 src/config/static/models.py 。" />
|
<a-alert type="warning" message="该提供商暂未适配获取模型列表的方法,如需添加模型,请编辑 src/config/static/models.py 。">
|
||||||
<img src="@/assets/pics/guides/how-to-add-models.png" alt="添加模型指引" style="width: 100%; height: 100%; margin-top: 16px;">
|
<template #description>
|
||||||
|
<a href="https://xerrors.github.io/Yuxi-Know/latest/intro/model-config.html" target="_blank">模型配置</a>
|
||||||
|
</template>
|
||||||
|
</a-alert>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user