commit
5fd20b9bbd
@ -16,4 +16,5 @@ RUN apt-get update && apt-get install ffmpeg libsm6 libxext6 -y
|
|||||||
|
|
||||||
# 复制代码到容器中
|
# 复制代码到容器中
|
||||||
COPY ../src /app/src
|
COPY ../src /app/src
|
||||||
|
COPY ../server /app/server
|
||||||
|
|
||||||
|
|||||||
@ -6,6 +6,7 @@ services:
|
|||||||
container_name: api-dev
|
container_name: api-dev
|
||||||
working_dir: /app
|
working_dir: /app
|
||||||
volumes:
|
volumes:
|
||||||
|
- ../server:/app/server
|
||||||
- ../src:/app/src
|
- ../src:/app/src
|
||||||
- ../saves_dev:/app/saves
|
- ../saves_dev:/app/saves
|
||||||
# - ${MODEL_DIR}:/models # 如果出现 undefined volume MODEL_DIR: invalid compose project,请添加此环境变量或者注释此行
|
# - ${MODEL_DIR}:/models # 如果出现 undefined volume MODEL_DIR: invalid compose project,请添加此环境变量或者注释此行
|
||||||
|
|||||||
40
server/db_manager.py
Normal file
40
server/db_manager.py
Normal file
@ -0,0 +1,40 @@
|
|||||||
|
import os
|
||||||
|
import sqlite3
|
||||||
|
import pathlib
|
||||||
|
from sqlalchemy import create_engine
|
||||||
|
from sqlalchemy.orm import sessionmaker
|
||||||
|
from sqlalchemy.ext.declarative import declarative_base
|
||||||
|
|
||||||
|
from server.models.token_model import Base, AgentToken
|
||||||
|
|
||||||
|
class DBManager:
|
||||||
|
"""数据库管理器"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.db_path = os.path.join("saves", "data", "server.db")
|
||||||
|
self.ensure_db_dir()
|
||||||
|
|
||||||
|
# 创建SQLAlchemy引擎
|
||||||
|
self.engine = create_engine(f"sqlite:///{self.db_path}")
|
||||||
|
|
||||||
|
# 创建会话工厂
|
||||||
|
self.Session = sessionmaker(bind=self.engine)
|
||||||
|
|
||||||
|
# 确保表存在
|
||||||
|
self.create_tables()
|
||||||
|
|
||||||
|
def ensure_db_dir(self):
|
||||||
|
"""确保数据库目录存在"""
|
||||||
|
db_dir = os.path.dirname(self.db_path)
|
||||||
|
pathlib.Path(db_dir).mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
def create_tables(self):
|
||||||
|
"""创建数据库表"""
|
||||||
|
Base.metadata.create_all(self.engine)
|
||||||
|
|
||||||
|
def get_session(self):
|
||||||
|
"""获取数据库会话"""
|
||||||
|
return self.Session()
|
||||||
|
|
||||||
|
# 创建全局数据库管理器实例
|
||||||
|
db_manager = DBManager()
|
||||||
@ -2,7 +2,7 @@ import uvicorn
|
|||||||
|
|
||||||
from fastapi import FastAPI
|
from fastapi import FastAPI
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from src.routers import router
|
from server.routers import router
|
||||||
from src.utils.logging_config import logger
|
from src.utils.logging_config import logger
|
||||||
|
|
||||||
|
|
||||||
24
server/models/token_model.py
Normal file
24
server/models/token_model.py
Normal file
@ -0,0 +1,24 @@
|
|||||||
|
from sqlalchemy import Column, Integer, String, DateTime, ForeignKey
|
||||||
|
from sqlalchemy.ext.declarative import declarative_base
|
||||||
|
from sqlalchemy.sql import func
|
||||||
|
|
||||||
|
Base = declarative_base()
|
||||||
|
|
||||||
|
class AgentToken(Base):
|
||||||
|
"""智能体访问令牌模型"""
|
||||||
|
__tablename__ = 'agent_tokens'
|
||||||
|
|
||||||
|
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||||
|
agent_id = Column(String, nullable=False, index=True) # 智能体ID
|
||||||
|
name = Column(String, nullable=False) # 令牌名称
|
||||||
|
token = Column(String, nullable=False, unique=True) # 令牌值
|
||||||
|
created_at = Column(DateTime, default=func.now()) # 创建时间
|
||||||
|
|
||||||
|
def to_dict(self):
|
||||||
|
return {
|
||||||
|
"id": self.id,
|
||||||
|
"agent_id": self.agent_id,
|
||||||
|
"name": self.name,
|
||||||
|
"token": self.token,
|
||||||
|
"created_at": self.created_at.isoformat() if self.created_at else None
|
||||||
|
}
|
||||||
13
server/routers/__init__.py
Normal file
13
server/routers/__init__.py
Normal file
@ -0,0 +1,13 @@
|
|||||||
|
from fastapi import APIRouter
|
||||||
|
from server.routers.chat_router import chat
|
||||||
|
from server.routers.data_router import data
|
||||||
|
from server.routers.base_router import base
|
||||||
|
from server.routers.tool_router import tool
|
||||||
|
from server.routers.admin_router import admin
|
||||||
|
|
||||||
|
router = APIRouter()
|
||||||
|
router.include_router(base)
|
||||||
|
router.include_router(chat)
|
||||||
|
router.include_router(data)
|
||||||
|
router.include_router(tool)
|
||||||
|
router.include_router(admin)
|
||||||
103
server/routers/admin_router.py
Normal file
103
server/routers/admin_router.py
Normal file
@ -0,0 +1,103 @@
|
|||||||
|
import secrets
|
||||||
|
import string
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||||
|
from pydantic import BaseModel
|
||||||
|
from typing import List, Optional
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from server.db_manager import db_manager
|
||||||
|
from server.models.token_model import AgentToken
|
||||||
|
|
||||||
|
admin = APIRouter(prefix="/admin", tags=["admin"])
|
||||||
|
|
||||||
|
# 依赖项:获取数据库会话
|
||||||
|
def get_db():
|
||||||
|
db = db_manager.get_session()
|
||||||
|
try:
|
||||||
|
yield db
|
||||||
|
finally:
|
||||||
|
db.close()
|
||||||
|
|
||||||
|
# 请求和响应模型
|
||||||
|
class TokenCreate(BaseModel):
|
||||||
|
agent_id: str
|
||||||
|
name: str
|
||||||
|
|
||||||
|
class TokenVerify(BaseModel):
|
||||||
|
agent_id: str
|
||||||
|
token: str
|
||||||
|
|
||||||
|
class TokenResponse(BaseModel):
|
||||||
|
id: int
|
||||||
|
agent_id: str
|
||||||
|
name: str
|
||||||
|
token: str
|
||||||
|
created_at: str
|
||||||
|
|
||||||
|
# 生成随机token
|
||||||
|
def generate_token(length=32):
|
||||||
|
alphabet = string.ascii_letters + string.digits
|
||||||
|
return ''.join(secrets.choice(alphabet) for _ in range(length))
|
||||||
|
|
||||||
|
@admin.get("/tokens", response_model=List[TokenResponse])
|
||||||
|
async def get_agent_tokens(
|
||||||
|
agent_id: Optional[str] = Query(None),
|
||||||
|
db: Session = Depends(get_db)
|
||||||
|
):
|
||||||
|
"""获取智能体的token列表"""
|
||||||
|
query = db.query(AgentToken)
|
||||||
|
if agent_id:
|
||||||
|
query = query.filter(AgentToken.agent_id == agent_id)
|
||||||
|
tokens = query.all()
|
||||||
|
return [token.to_dict() for token in tokens]
|
||||||
|
|
||||||
|
@admin.post("/tokens", response_model=TokenResponse)
|
||||||
|
async def create_token(
|
||||||
|
token_data: TokenCreate,
|
||||||
|
db: Session = Depends(get_db)
|
||||||
|
):
|
||||||
|
"""创建新的token"""
|
||||||
|
# 生成随机token
|
||||||
|
token_value = generate_token()
|
||||||
|
|
||||||
|
# 创建token记录
|
||||||
|
new_token = AgentToken(
|
||||||
|
agent_id=token_data.agent_id,
|
||||||
|
name=token_data.name,
|
||||||
|
token=token_value
|
||||||
|
)
|
||||||
|
|
||||||
|
# 保存到数据库
|
||||||
|
db.add(new_token)
|
||||||
|
db.commit()
|
||||||
|
db.refresh(new_token)
|
||||||
|
|
||||||
|
return new_token.to_dict()
|
||||||
|
|
||||||
|
@admin.delete("/tokens/{token_id}", response_model=dict)
|
||||||
|
async def delete_token(token_id: int, db: Session = Depends(get_db)):
|
||||||
|
"""删除token"""
|
||||||
|
token = db.query(AgentToken).filter(AgentToken.id == token_id).first()
|
||||||
|
if not token:
|
||||||
|
raise HTTPException(status_code=404, detail="Token not found")
|
||||||
|
|
||||||
|
db.delete(token)
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
return {"success": True, "message": "Token deleted"}
|
||||||
|
|
||||||
|
@admin.post("/verify_token")
|
||||||
|
async def verify_agent_token(
|
||||||
|
token_data: TokenVerify,
|
||||||
|
db: Session = Depends(get_db)
|
||||||
|
):
|
||||||
|
"""验证智能体访问令牌"""
|
||||||
|
token = db.query(AgentToken).filter(
|
||||||
|
AgentToken.agent_id == token_data.agent_id,
|
||||||
|
AgentToken.token == token_data.token
|
||||||
|
).first()
|
||||||
|
|
||||||
|
if not token:
|
||||||
|
raise HTTPException(status_code=401, detail="Invalid token")
|
||||||
|
|
||||||
|
return {"success": True, "message": "Token verified"}
|
||||||
@ -1,8 +1,9 @@
|
|||||||
|
import os
|
||||||
import json
|
import json
|
||||||
import asyncio
|
import asyncio
|
||||||
import traceback
|
import traceback
|
||||||
import uuid
|
import uuid
|
||||||
from fastapi import APIRouter, Body
|
from fastapi import APIRouter, Body, Depends, HTTPException
|
||||||
from fastapi.responses import StreamingResponse
|
from fastapi.responses import StreamingResponse
|
||||||
from langchain_core.messages import AIMessageChunk
|
from langchain_core.messages import AIMessageChunk
|
||||||
|
|
||||||
@ -14,8 +15,6 @@ from src.utils.logging_config import logger
|
|||||||
|
|
||||||
chat = APIRouter(prefix="/chat")
|
chat = APIRouter(prefix="/chat")
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@chat.get("/")
|
@chat.get("/")
|
||||||
async def chat_get():
|
async def chat_get():
|
||||||
return "Chat Get!"
|
return "Chat Get!"
|
||||||
@ -123,60 +122,77 @@ async def call(query: str = Body(...), meta: dict = Body(None)):
|
|||||||
|
|
||||||
@chat.get("/agent")
|
@chat.get("/agent")
|
||||||
async def get_agent():
|
async def get_agent():
|
||||||
agents = [{
|
agents = [agent.get_info() for agent in agent_manager.agents.values()]
|
||||||
"name": agent.name,
|
|
||||||
"description": agent.description,
|
|
||||||
"config_schema": agent.config_schema
|
|
||||||
} for agent in agent_manager.agents.values()]
|
|
||||||
return {"agents": agents}
|
return {"agents": agents}
|
||||||
|
|
||||||
@chat.post("/agent/{agent_name}")
|
@chat.post("/agent/{agent_name}")
|
||||||
def chat_agent(agent_name: str,
|
def chat_agent(agent_name: str,
|
||||||
query: str = Body(...),
|
query: str = Body(...),
|
||||||
meta: dict = Body({}),
|
|
||||||
history: list = Body(...),
|
history: list = Body(...),
|
||||||
thread_id: str | None = Body(None)):
|
config: dict = Body({}),
|
||||||
|
meta: dict = Body({})):
|
||||||
|
|
||||||
meta["server_model_name"] = agent_name
|
meta.update({
|
||||||
agent = agent_manager.get_runnable_agent(agent_name)
|
"query": query,
|
||||||
|
"agent_name": agent_name,
|
||||||
history_manager = HistoryManager(history)
|
"server_model_name": config.get("model", agent_name) ,
|
||||||
messages = history_manager.get_history_with_msg(query, max_rounds=meta.get('history_round'))
|
"thread_id": config.get("thread_id"),
|
||||||
history_manager.add_user(query) # 注意这里使用原始查询
|
})
|
||||||
|
|
||||||
runnable_config = {
|
|
||||||
"configurable": {
|
|
||||||
"thread_id": thread_id or str(uuid.uuid4()),
|
|
||||||
"use_web": meta.get("use_web", False),
|
|
||||||
"return_keys": []
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
|
# 将meta和thread_id整合到config中
|
||||||
def make_chunk(content=None, **kwargs):
|
def make_chunk(content=None, **kwargs):
|
||||||
|
|
||||||
return json.dumps({
|
return json.dumps({
|
||||||
|
"request_id": meta.get("request_id"),
|
||||||
"response": content,
|
"response": content,
|
||||||
"model_name": agent_name,
|
|
||||||
"meta": meta,
|
|
||||||
**kwargs
|
**kwargs
|
||||||
}, ensure_ascii=False).encode('utf-8') + b"\n"
|
}, ensure_ascii=False).encode('utf-8') + b"\n"
|
||||||
|
|
||||||
|
try:
|
||||||
|
agent = agent_manager.get_runnable_agent(agent_name)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting agent {agent_name}: {e}")
|
||||||
|
return StreamingResponse(make_chunk(message=f"Error getting agent {agent_name}: {e}", status="error"), media_type='application/json')
|
||||||
|
|
||||||
|
# 从config中获取history_round
|
||||||
|
history_round = config.get("history_round")
|
||||||
|
history_manager = HistoryManager(history)
|
||||||
|
messages = history_manager.get_history_with_msg(query, max_rounds=history_round)
|
||||||
|
history_manager.add_user(query)
|
||||||
|
|
||||||
|
# 如果没有thread_id则生成一个
|
||||||
|
if "thread_id" not in config or not config["thread_id"]:
|
||||||
|
config["thread_id"] = str(uuid.uuid4())
|
||||||
|
|
||||||
|
# 构造运行时配置
|
||||||
|
runnable_config = {
|
||||||
|
"configurable": {
|
||||||
|
**config
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
def stream_messages():
|
def stream_messages():
|
||||||
content = ""
|
content = ""
|
||||||
yield make_chunk(status="waiting")
|
yield make_chunk(status="init", meta=meta)
|
||||||
for msg, metadata in agent.stream_messages(messages, runnable_config):
|
try:
|
||||||
# logger.debug(f">>>>> msg: {msg.model_dump()}, >>>>>>> {metadata=}")
|
for msg, metadata in agent.stream_messages(messages, config_schema=runnable_config):
|
||||||
if isinstance(msg, AIMessageChunk) and msg.content != "<tool_call>":
|
if isinstance(msg, AIMessageChunk) and msg.content != "<tool_call>":
|
||||||
content += msg.content
|
content += msg.content
|
||||||
yield make_chunk(content=msg.content,
|
yield make_chunk(content=msg.content,
|
||||||
msg=msg.model_dump(),
|
msg=msg.model_dump(),
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
status="loading")
|
status="loading")
|
||||||
else:
|
else:
|
||||||
yield make_chunk(msg=msg.model_dump(),
|
yield make_chunk(msg=msg.model_dump(),
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
status="loading")
|
status="loading")
|
||||||
|
|
||||||
yield make_chunk(status="finished", history=history_manager.update_ai(content))
|
yield make_chunk(status="finished",
|
||||||
|
history=history_manager.update_ai(content),
|
||||||
|
meta=meta)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error streaming messages: {e}")
|
||||||
|
yield make_chunk(message=f"Error streaming messages: {e}", status="error")
|
||||||
|
|
||||||
return StreamingResponse(stream_messages(), media_type='application/json')
|
return StreamingResponse(stream_messages(), media_type='application/json')
|
||||||
|
|
||||||
@ -72,7 +72,7 @@ async def create_document_by_file(db_id: str = Body(...), files: List[str] = Bod
|
|||||||
|
|
||||||
@data.post("/add-by-chunks")
|
@data.post("/add-by-chunks")
|
||||||
async def add_by_chunks(db_id: str = Body(...), file_chunks: dict = Body(...)):
|
async def add_by_chunks(db_id: str = Body(...), file_chunks: dict = Body(...)):
|
||||||
logger.debug(f"Add chunks in {db_id}: {file_chunks}")
|
# logger.debug(f"Add chunks in {db_id}: {len(file_chunks)} chunks")
|
||||||
try:
|
try:
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_event_loop()
|
||||||
await loop.run_in_executor(
|
await loop.run_in_executor(
|
||||||
@ -50,7 +50,7 @@ async def route_index():
|
|||||||
description=agent.description,
|
description=agent.description,
|
||||||
url=f"/agent/{agent.name}",
|
url=f"/agent/{agent.name}",
|
||||||
method="POST",
|
method="POST",
|
||||||
metadata=agent.config_schema,
|
metadata=agent.config_schema.to_dict(),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -1,36 +1,36 @@
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import datetime, timezone
|
|
||||||
|
|
||||||
from langchain_core.tools import tool
|
|
||||||
from langchain_openai import ChatOpenAI
|
|
||||||
|
|
||||||
from src.agents.registry import Configuration
|
from src.agents.registry import Configuration
|
||||||
|
|
||||||
def get_default_requirements():
|
|
||||||
return ["TAVILY_API_KEY"]
|
|
||||||
|
|
||||||
@tool
|
|
||||||
def multiply(first_int: int, second_int: int) -> int:
|
|
||||||
"""Multiply two integers together."""
|
|
||||||
return first_int * second_int
|
|
||||||
|
|
||||||
@dataclass(kw_only=True)
|
@dataclass(kw_only=True)
|
||||||
class ChatbotConfiguration(Configuration):
|
class ChatbotConfiguration(Configuration):
|
||||||
"""Chatbot 的配置"""
|
"""Chatbot 的配置
|
||||||
|
|
||||||
|
配置说明:
|
||||||
|
|
||||||
|
metadata 中 configurable 为 True 的配置项可以被用户配置,
|
||||||
|
configurable 为 False 的配置项不能被用户配置,只能由开发者预设。
|
||||||
|
"""
|
||||||
|
|
||||||
system_prompt: str = field(
|
system_prompt: str = field(
|
||||||
default=f"You are a helpful assistant. Now is {datetime.now(tz=timezone.utc).isoformat()}",
|
default="You are a helpful assistant.",
|
||||||
metadata={
|
metadata={
|
||||||
"description": "The system prompt to use for the agent's interactions. "
|
"name": "系统提示词",
|
||||||
"This prompt sets the context and behavior for the agent."
|
"configurable": True,
|
||||||
|
"description": "用来描述智能体的角色和行为"
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
model: str = field(
|
model: str = field(
|
||||||
default="zhipu/glm-4-plus",
|
default="zhipu/glm-4-plus",
|
||||||
metadata={
|
metadata={
|
||||||
"description": "The name of the language model to use for the agent's main interactions. "
|
"name": "智能体模型",
|
||||||
"Should be in the form: provider/model-name."
|
"configurable": True,
|
||||||
|
"options": [
|
||||||
|
"zhipu/glm-4-plus",
|
||||||
|
"siliconflow/Qwen/QwQ-32B",
|
||||||
|
"siliconflow/deepseek-ai/DeepSeek-V3",
|
||||||
|
],
|
||||||
|
"description": "智能体的驱动模型"
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -1,48 +1,62 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from datetime import datetime
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
from langchain_core.runnables import RunnableConfig
|
from langchain_core.runnables import RunnableConfig
|
||||||
from langgraph.graph import StateGraph, START, END
|
from langgraph.graph import StateGraph, START, END
|
||||||
from langgraph.prebuilt import ToolNode, tools_condition
|
from langgraph.prebuilt import ToolNode, tools_condition
|
||||||
from langgraph.checkpoint.memory import MemorySaver
|
from langgraph.checkpoint.memory import MemorySaver # 实际上没有起作用
|
||||||
from langchain_community.tools.tavily_search import TavilySearchResults
|
|
||||||
|
|
||||||
|
|
||||||
|
from src.utils import logger
|
||||||
from src.agents.registry import State, BaseAgent
|
from src.agents.registry import State, BaseAgent
|
||||||
from src.agents.utils import load_chat_model
|
from src.agents.utils import load_chat_model, get_cur_time_with_utc
|
||||||
from src.agents.chatbot.configuration import ChatbotConfiguration, multiply
|
from src.agents.chatbot.configuration import ChatbotConfiguration
|
||||||
|
from src.agents.tools_factory import _TOOLS_REGISTRY
|
||||||
|
|
||||||
class ChatbotAgent(BaseAgent):
|
class ChatbotAgent(BaseAgent):
|
||||||
name = "chatbot"
|
name = "chatbot"
|
||||||
description = "A chatbot that can answer questions and help with tasks."
|
description = "A chatbot that can answer questions and help with tasks."
|
||||||
_graph_cache = None
|
requirements = ["TAVILY_API_KEY", "ZHIPUAI_API_KEY"]
|
||||||
config_schema = ChatbotConfiguration.to_dict()
|
all_tools = ["TavilySearchResults", "multiply", "add", "subtract", "divide"]
|
||||||
|
config_schema = ChatbotConfiguration
|
||||||
|
|
||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
def _get_tools(self, config_schema: RunnableConfig):
|
def _get_tools(self, config_schema: RunnableConfig):
|
||||||
"""根据配置获取工具"""
|
"""根据配置获取工具,如果配置为空,则使用所有工具,如果配置为列表,则使用列表中的工具,
|
||||||
tools = [multiply, TavilySearchResults(max_results=10)]
|
如果配置为其他类型,则抛出错误"""
|
||||||
return tools
|
conf_tools = config_schema.get("tools")
|
||||||
|
if conf_tools == None:
|
||||||
|
tool_names = self.all_tools
|
||||||
|
elif isinstance(conf_tools, list):
|
||||||
|
tool_names = [tool for tool in self.all_tools if tool in conf_tools]
|
||||||
|
else:
|
||||||
|
raise ValueError(f"tools 配置错误: {conf_tools}")
|
||||||
|
|
||||||
|
logger.info(f"Tools: {tool_names}")
|
||||||
|
return [_TOOLS_REGISTRY[tool] for tool in tool_names]
|
||||||
|
|
||||||
def llm_call(self, state: State, config: RunnableConfig = None) -> dict[str, Any]:
|
def llm_call(self, state: State, config: RunnableConfig = None) -> dict[str, Any]:
|
||||||
"""调用 llm 模型"""
|
"""调用 llm 模型"""
|
||||||
config_schema = config or {}
|
config_schema = config or {}
|
||||||
conf = ChatbotConfiguration.from_runnable_config(config_schema)
|
conf = self.config_schema.from_runnable_config(config_schema)
|
||||||
|
|
||||||
|
system_prompt = f"{conf.system_prompt} Now is {get_cur_time_with_utc()}"
|
||||||
model = load_chat_model(conf.model)
|
model = load_chat_model(conf.model)
|
||||||
model_with_tools = model.bind_tools(self._get_tools(config_schema))
|
model_with_tools = model.bind_tools(self._get_tools(config_schema))
|
||||||
|
logger.info(f"llm_call with config: {conf}, {conf.model}")
|
||||||
|
|
||||||
res = model_with_tools.invoke(
|
res = model_with_tools.invoke(
|
||||||
[{"role": "system", "content": conf.system_prompt}, *state["messages"]]
|
[{"role": "system", "content": system_prompt}, *state["messages"]]
|
||||||
)
|
)
|
||||||
return {"messages": [res]}
|
return {"messages": [res]}
|
||||||
|
|
||||||
def get_graph(self, config_schema: RunnableConfig = None):
|
def get_graph(self, config_schema: RunnableConfig = None, **kwargs):
|
||||||
"""构建图"""
|
"""构建图"""
|
||||||
workflow = StateGraph(State, config_schema=ChatbotConfiguration)
|
workflow = StateGraph(State, config_schema=self.config_schema)
|
||||||
workflow.add_node("chatbot", self.llm_call)
|
workflow.add_node("chatbot", self.llm_call)
|
||||||
workflow.add_node("tools", ToolNode(tools=self._get_tools(config_schema)))
|
workflow.add_node("tools", ToolNode(tools=self._get_tools(config_schema)))
|
||||||
workflow.add_edge(START, "chatbot")
|
workflow.add_edge(START, "chatbot")
|
||||||
@ -56,20 +70,6 @@ class ChatbotAgent(BaseAgent):
|
|||||||
graph = workflow.compile(checkpointer=MemorySaver())
|
graph = workflow.compile(checkpointer=MemorySaver())
|
||||||
return graph
|
return graph
|
||||||
|
|
||||||
def stream_values(self, messages: list[str], config_schema: RunnableConfig = None):
|
|
||||||
graph = self.get_graph(config_schema)
|
|
||||||
for event in graph.stream({"messages": messages}, stream_mode="values", config=config_schema):
|
|
||||||
yield event["messages"]
|
|
||||||
|
|
||||||
def stream_messages(self, messages: list[str], config_schema: RunnableConfig = None):
|
|
||||||
graph = self.get_graph(config_schema)
|
|
||||||
conf = ChatbotConfiguration.from_runnable_config(config_schema)
|
|
||||||
for msg, metadata in graph.stream({"messages": messages}, stream_mode="messages", config=config_schema):
|
|
||||||
msg_type = msg.type
|
|
||||||
|
|
||||||
return_keys =conf.return_keys
|
|
||||||
if not return_keys or msg_type in return_keys:
|
|
||||||
yield msg, metadata
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
agent = ChatbotAgent(ChatbotConfiguration())
|
agent = ChatbotAgent(ChatbotConfiguration())
|
||||||
|
|||||||
@ -17,60 +17,12 @@ from src.agents.react.configuration import ReActConfiguration, multiply
|
|||||||
class ReActAgent(BaseAgent):
|
class ReActAgent(BaseAgent):
|
||||||
name = "react"
|
name = "react"
|
||||||
description = "A react agent that can answer questions and help with tasks."
|
description = "A react agent that can answer questions and help with tasks."
|
||||||
_graph_cache = None
|
|
||||||
config_schema = ReActConfiguration.to_dict()
|
|
||||||
|
|
||||||
def __init__(self, **kwargs):
|
def get_graph(self, **kwargs):
|
||||||
super().__init__(**kwargs)
|
|
||||||
|
|
||||||
def _get_tools(self, config_schema: RunnableConfig):
|
|
||||||
"""根据配置获取工具"""
|
|
||||||
tools = [multiply, TavilySearchResults(max_results=10)]
|
|
||||||
return tools
|
|
||||||
|
|
||||||
def llm_call(self, state: State, config: RunnableConfig = None) -> dict[str, Any]:
|
|
||||||
"""调用 llm 模型"""
|
|
||||||
config_schema = config or {}
|
|
||||||
conf = ReActConfiguration.from_runnable_config(config_schema)
|
|
||||||
model = load_chat_model(conf.model)
|
|
||||||
model_with_tools = model.bind_tools(self._get_tools(config_schema))
|
|
||||||
|
|
||||||
res = model_with_tools.invoke(
|
|
||||||
[{"role": "system", "content": conf.system_prompt}, *state["messages"]]
|
|
||||||
)
|
|
||||||
return {"messages": [res]}
|
|
||||||
|
|
||||||
def get_graph(self, config_schema: RunnableConfig = None):
|
|
||||||
"""构建图"""
|
"""构建图"""
|
||||||
workflow = StateGraph(State, config_schema=ReActConfiguration)
|
from .workflows import graph
|
||||||
workflow.add_node("react", self.llm_call)
|
|
||||||
workflow.add_node("tools", ToolNode(tools=self._get_tools(config_schema)))
|
|
||||||
workflow.add_edge(START, "react")
|
|
||||||
workflow.add_conditional_edges(
|
|
||||||
"react",
|
|
||||||
tools_condition,
|
|
||||||
)
|
|
||||||
workflow.add_edge("tools", "react")
|
|
||||||
workflow.add_edge("react", END)
|
|
||||||
|
|
||||||
graph = workflow.compile(checkpointer=MemorySaver())
|
|
||||||
return graph
|
return graph
|
||||||
|
|
||||||
def stream_values(self, messages: list[str], config_schema: RunnableConfig = None):
|
|
||||||
graph = self.get_graph(config_schema)
|
|
||||||
for event in graph.stream({"messages": messages}, stream_mode="values", config=config_schema):
|
|
||||||
yield event["messages"]
|
|
||||||
|
|
||||||
def stream_messages(self, messages: list[str], config_schema: RunnableConfig = None):
|
|
||||||
graph = self.get_graph(config_schema)
|
|
||||||
conf = ReActConfiguration.from_runnable_config(config_schema)
|
|
||||||
for msg, metadata in graph.stream({"messages": messages}, stream_mode="messages", config=config_schema):
|
|
||||||
msg_type = msg.type
|
|
||||||
|
|
||||||
return_keys =conf.return_keys
|
|
||||||
if not return_keys or msg_type in return_keys:
|
|
||||||
yield msg, metadata
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
agent = ReActAgent(ReActConfiguration())
|
agent = ReActAgent(ReActConfiguration())
|
||||||
|
|
||||||
|
|||||||
36
src/agents/react/workflows.py
Normal file
36
src/agents/react/workflows.py
Normal file
@ -0,0 +1,36 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
from langchain_openai import ChatOpenAI
|
||||||
|
|
||||||
|
model = ChatOpenAI(model="glm-4-plus",
|
||||||
|
api_key=os.getenv("ZHIPUAI_API_KEY"),
|
||||||
|
base_url="https://open.bigmodel.cn/api/paas/v4/",
|
||||||
|
temperature=0)
|
||||||
|
|
||||||
|
|
||||||
|
# For this tutorial we will use custom tool that returns pre-defined values for weather in two cities (NYC & SF)
|
||||||
|
|
||||||
|
from typing import Literal
|
||||||
|
|
||||||
|
from langchain_core.tools import tool
|
||||||
|
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def get_weather(city: Literal["nyc", "sf"]):
|
||||||
|
"""Use this to get weather information."""
|
||||||
|
if city == "nyc":
|
||||||
|
return "It might be cloudy in nyc"
|
||||||
|
elif city == "sf":
|
||||||
|
return "It's always sunny in sf"
|
||||||
|
else:
|
||||||
|
raise AssertionError("Unknown city")
|
||||||
|
|
||||||
|
|
||||||
|
tools = [get_weather]
|
||||||
|
|
||||||
|
|
||||||
|
# Define the graph
|
||||||
|
|
||||||
|
from langgraph.prebuilt import create_react_agent
|
||||||
|
|
||||||
|
graph = create_react_agent(model, tools=tools)
|
||||||
@ -1,5 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
|
||||||
from typing import Type, Annotated, Optional, TypedDict
|
from typing import Type, Annotated, Optional, TypedDict
|
||||||
from abc import abstractmethod
|
from abc import abstractmethod
|
||||||
from dataclasses import dataclass, fields, field
|
from dataclasses import dataclass, fields, field
|
||||||
@ -10,7 +12,7 @@ from langgraph.graph.state import CompiledStateGraph
|
|||||||
from langgraph.graph.message import add_messages
|
from langgraph.graph.message import add_messages
|
||||||
|
|
||||||
from src.config import SimpleConfig
|
from src.config import SimpleConfig
|
||||||
|
from src.utils import logger
|
||||||
|
|
||||||
class State(TypedDict):
|
class State(TypedDict):
|
||||||
"""
|
"""
|
||||||
@ -24,7 +26,7 @@ class State(TypedDict):
|
|||||||
|
|
||||||
|
|
||||||
@dataclass(kw_only=True)
|
@dataclass(kw_only=True)
|
||||||
class Configuration(SimpleConfig):
|
class Configuration(dict):
|
||||||
"""
|
"""
|
||||||
定义一个基础 Configuration 供 各类 graph 继承
|
定义一个基础 Configuration 供 各类 graph 继承
|
||||||
"""
|
"""
|
||||||
@ -40,15 +42,73 @@ class Configuration(SimpleConfig):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def to_dict(cls):
|
def to_dict(cls):
|
||||||
return {f.name: getattr(cls, f.name) for f in fields(cls) if f.init}
|
# 创建一个实例来处理 default_factory
|
||||||
|
instance = cls()
|
||||||
|
confs = {}
|
||||||
|
configurable_items = {}
|
||||||
|
for f in fields(cls):
|
||||||
|
if f.init and not f.metadata.get("hide", False):
|
||||||
|
value = getattr(instance, f.name)
|
||||||
|
if callable(value) and hasattr(value, "__call__"):
|
||||||
|
confs[f.name] = value()
|
||||||
|
else:
|
||||||
|
confs[f.name] = value
|
||||||
|
|
||||||
|
if f.metadata.get("configurable"):
|
||||||
|
configurable_items[f.name] = {
|
||||||
|
"type": f.type.__name__,
|
||||||
|
"name": f.metadata.get("name", f.name),
|
||||||
|
"options": f.metadata.get("options", []),
|
||||||
|
"default": f.default,
|
||||||
|
"description": f.metadata.get("description", ""),
|
||||||
|
}
|
||||||
|
confs["configurable_items"] = configurable_items
|
||||||
|
return confs
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class BaseAgent():
|
class BaseAgent():
|
||||||
|
|
||||||
|
"""
|
||||||
|
定义一个基础 Agent 供 各类 graph 继承
|
||||||
|
"""
|
||||||
|
|
||||||
|
name: str = field(default="base_agent")
|
||||||
|
description: str = field(default="base_agent")
|
||||||
|
config_schema: Configuration = Configuration
|
||||||
|
requirements: list[str]
|
||||||
|
|
||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
pass
|
self.check_requirements()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_info(cls):
|
||||||
|
return {
|
||||||
|
"name": cls.name,
|
||||||
|
"description": cls.description,
|
||||||
|
"config_schema": cls.config_schema.to_dict(),
|
||||||
|
"requirements": cls.requirements if hasattr(cls, "requirements") else [],
|
||||||
|
"all_tools": cls.all_tools if hasattr(cls, "all_tools") else [],
|
||||||
|
}
|
||||||
|
|
||||||
|
def check_requirements(self):
|
||||||
|
if not hasattr(self, "requirements") or not self.requirements:
|
||||||
|
return
|
||||||
|
for requirement in self.requirements:
|
||||||
|
if requirement not in os.environ:
|
||||||
|
raise ValueError(f"{requirement} is not set")
|
||||||
|
|
||||||
|
def stream_values(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs):
|
||||||
|
graph = self.get_graph(config_schema=config_schema, **kwargs)
|
||||||
|
for event in graph.stream({"messages": messages}, stream_mode="values", config=config_schema):
|
||||||
|
yield event["messages"]
|
||||||
|
|
||||||
|
def stream_messages(self, messages: list[str], config_schema: RunnableConfig = None, **kwargs):
|
||||||
|
graph = self.get_graph(config_schema=config_schema, **kwargs)
|
||||||
|
|
||||||
|
for msg, metadata in graph.stream({"messages": messages}, stream_mode="messages", config=config_schema):
|
||||||
|
yield msg, metadata
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def get_graph(self) -> CompiledStateGraph:
|
def get_graph(self, **kwargs) -> CompiledStateGraph:
|
||||||
pass
|
pass
|
||||||
@ -5,11 +5,8 @@ from typing import Any, Callable, Optional, Type, Union
|
|||||||
|
|
||||||
|
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
from langchain.agents import tool
|
from langchain_core.tools import tool, BaseTool
|
||||||
from langchain_core.tools import Tool, BaseTool
|
from langchain_community.tools.tavily_search import TavilySearchResults
|
||||||
|
|
||||||
|
|
||||||
_TOOLS_REGISTRY = {}
|
|
||||||
|
|
||||||
# refs https://github.com/chatchat-space/LangGraph-Chatchat chatchat-server/chatchat/server/agent/tools_factory/tools_registry.py
|
# refs https://github.com/chatchat-space/LangGraph-Chatchat chatchat-server/chatchat/server/agent/tools_factory/tools_registry.py
|
||||||
def regist_tool(
|
def regist_tool(
|
||||||
@ -95,3 +92,33 @@ class BaseToolOutput:
|
|||||||
else:
|
else:
|
||||||
return str(self.data)
|
return str(self.data)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def multiply(first_int: int, second_int: int) -> int:
|
||||||
|
"""Multiply two integers together."""
|
||||||
|
return first_int * second_int
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def add(first_int: int, second_int: int) -> int:
|
||||||
|
"""Add two integers together."""
|
||||||
|
return first_int + second_int
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def subtract(first_int: int, second_int: int) -> int:
|
||||||
|
"""Subtract two integers."""
|
||||||
|
return first_int - second_int
|
||||||
|
|
||||||
|
@tool
|
||||||
|
def divide(first_int: int, second_int: int) -> int:
|
||||||
|
"""Divide two integers."""
|
||||||
|
return first_int / second_int
|
||||||
|
|
||||||
|
|
||||||
|
_TOOLS_REGISTRY = {
|
||||||
|
"multiply": multiply,
|
||||||
|
"add": add,
|
||||||
|
"subtract": subtract,
|
||||||
|
"divide": divide,
|
||||||
|
"TavilySearchResults": TavilySearchResults(max_results=10),
|
||||||
|
}
|
||||||
|
|||||||
@ -1,3 +1,5 @@
|
|||||||
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
from src.models import select_model
|
from src.models import select_model
|
||||||
from src.agents.registry import BaseAgent
|
from src.agents.registry import BaseAgent
|
||||||
from langchain_core.language_models import BaseChatModel
|
from langchain_core.language_models import BaseChatModel
|
||||||
@ -7,14 +9,17 @@ from langchain_core.messages import AIMessageChunk, ToolMessage
|
|||||||
|
|
||||||
|
|
||||||
|
|
||||||
def load_chat_model(fully_specified_name: str) -> BaseChatModel:
|
def load_chat_model(fully_specified_name: str, **kwargs) -> BaseChatModel:
|
||||||
"""Load a chat model from a fully specified name.
|
"""Load a chat model from a fully specified name.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
fully_specified_name (str): String in the format 'provider/model'.
|
fully_specified_name (str): String in the format 'provider/model'.
|
||||||
|
**kwargs: Additional parameters to pass to the model.
|
||||||
"""
|
"""
|
||||||
provider, model = fully_specified_name.split("/", maxsplit=1)
|
provider, model = fully_specified_name.split("/", maxsplit=1)
|
||||||
return select_model(model_name=model, model_provider=provider).chat_open_ai
|
model_instance = select_model(model_name=model, model_provider=provider)
|
||||||
|
|
||||||
|
return model_instance.chat_open_ai
|
||||||
|
|
||||||
|
|
||||||
def agent_cli(agent: BaseAgent, config: RunnableConfig = None):
|
def agent_cli(agent: BaseAgent, config: RunnableConfig = None):
|
||||||
@ -49,3 +54,7 @@ def agent_cli(agent: BaseAgent, config: RunnableConfig = None):
|
|||||||
|
|
||||||
if isinstance(msg, ToolMessage):
|
if isinstance(msg, ToolMessage):
|
||||||
print(f"Tool: {msg.content}")
|
print(f"Tool: {msg.content}")
|
||||||
|
|
||||||
|
def get_cur_time_with_utc():
|
||||||
|
return datetime.now(tz=timezone.utc).isoformat()
|
||||||
|
|
||||||
|
|||||||
@ -28,7 +28,7 @@ class SimpleConfig(dict):
|
|||||||
return self.get(self.__key(key))
|
return self.get(self.__key(key))
|
||||||
|
|
||||||
def __getitem__(self, key):
|
def __getitem__(self, key):
|
||||||
return super().get(self.__key(key))
|
return self.get(self.__key(key))
|
||||||
|
|
||||||
def __setitem__(self, key, value):
|
def __setitem__(self, key, value):
|
||||||
return super().__setitem__(self.__key(key), value)
|
return super().__setitem__(self.__key(key), value)
|
||||||
|
|||||||
@ -231,7 +231,7 @@ class Retriever:
|
|||||||
return formatted_results
|
return formatted_results
|
||||||
|
|
||||||
def format_query_results(self, results):
|
def format_query_results(self, results):
|
||||||
logger.debug(f"Graph Query Results: {results}")
|
# logger.debug(f"Graph Query Results: {results}")
|
||||||
formatted_results = {"nodes": [], "edges": []}
|
formatted_results = {"nodes": [], "edges": []}
|
||||||
node_dict = {}
|
node_dict = {}
|
||||||
|
|
||||||
|
|||||||
@ -11,7 +11,8 @@ class OpenAIBase():
|
|||||||
self.model_name = model_name
|
self.model_name = model_name
|
||||||
self.chat_open_ai = ChatOpenAI(model=model_name,
|
self.chat_open_ai = ChatOpenAI(model=model_name,
|
||||||
api_key=api_key,
|
api_key=api_key,
|
||||||
base_url=base_url)
|
base_url=base_url,
|
||||||
|
temperature=0.7)
|
||||||
|
|
||||||
def predict(self, message, stream=False):
|
def predict(self, message, stream=False):
|
||||||
if isinstance(message, str):
|
if isinstance(message, str):
|
||||||
|
|||||||
@ -1,11 +0,0 @@
|
|||||||
from fastapi import APIRouter
|
|
||||||
from src.routers.chat_router import chat
|
|
||||||
from src.routers.data_router import data
|
|
||||||
from src.routers.base_router import base
|
|
||||||
from src.routers.tool_router import tool
|
|
||||||
|
|
||||||
router = APIRouter()
|
|
||||||
router.include_router(base)
|
|
||||||
router.include_router(chat)
|
|
||||||
router.include_router(data)
|
|
||||||
router.include_router(tool)
|
|
||||||
@ -10,12 +10,17 @@
|
|||||||
padding: 1rem;
|
padding: 1rem;
|
||||||
}
|
}
|
||||||
|
|
||||||
.message-md pre:has(code.hljs) {
|
.message-md pre:has(code) {
|
||||||
padding: 0;
|
padding: 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
.message-md pre code.hljs {
|
.message-md pre code {
|
||||||
font-size: 0.8rem;
|
font-size: 13px;
|
||||||
|
font-family: 'Menlo', 'Monaco', 'Consolas', 'PingFang SC', 'Microsoft YaHei', 'Hiragino Sans GB', 'Source Han Sans CN', 'Courier New', monospace;
|
||||||
|
line-height: 1.5;
|
||||||
|
letter-spacing: 0.025em;
|
||||||
|
tab-size: 4;
|
||||||
|
-moz-tab-size: 4;
|
||||||
background-color: var(--gray-100);
|
background-color: var(--gray-100);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
1272
web/src/components/AgentChatComponent.vue
Normal file
1272
web/src/components/AgentChatComponent.vue
Normal file
File diff suppressed because it is too large
Load Diff
@ -6,6 +6,9 @@
|
|||||||
<!-- 助手消息 -->
|
<!-- 助手消息 -->
|
||||||
<div v-else-if="message.role === 'assistant' || message.role === 'received'" class="assistant-message">
|
<div v-else-if="message.role === 'assistant' || message.role === 'received'" class="assistant-message">
|
||||||
<!-- 推理过程 (ChatComponent特有) -->
|
<!-- 推理过程 (ChatComponent特有) -->
|
||||||
|
<p v-if="debugMode">
|
||||||
|
{{ message.status }}
|
||||||
|
</p>
|
||||||
<div v-if="message.reasoning_content" class="reasoning-box">
|
<div v-if="message.reasoning_content" class="reasoning-box">
|
||||||
<a-collapse v-model:activeKey="reasoningActiveKey" :bordered="false">
|
<a-collapse v-model:activeKey="reasoningActiveKey" :bordered="false">
|
||||||
<template #expandIcon="{ isActive }">
|
<template #expandIcon="{ isActive }">
|
||||||
@ -50,7 +53,7 @@
|
|||||||
<slot name="tool-calls"></slot>
|
<slot name="tool-calls"></slot>
|
||||||
|
|
||||||
<div v-if="(message.role=='received' || message.role=='assistant') && message.status=='finished' && showRefs">
|
<div v-if="(message.role=='received' || message.role=='assistant') && message.status=='finished' && showRefs">
|
||||||
<RefsComponent :message="message" @retry="emit('retry')" />
|
<RefsComponent :message="message" :show-refs="showRefs" @retry="emit('retry')" />
|
||||||
</div>
|
</div>
|
||||||
<!-- 错误消息 -->
|
<!-- 错误消息 -->
|
||||||
</div>
|
</div>
|
||||||
@ -77,11 +80,6 @@ const props = defineProps({
|
|||||||
type: Object,
|
type: Object,
|
||||||
required: true
|
required: true
|
||||||
},
|
},
|
||||||
// 已渲染的HTML内容
|
|
||||||
contentHtml: {
|
|
||||||
type: String,
|
|
||||||
default: ''
|
|
||||||
},
|
|
||||||
// 是否正在处理中
|
// 是否正在处理中
|
||||||
isProcessing: {
|
isProcessing: {
|
||||||
type: Boolean,
|
type: Boolean,
|
||||||
@ -94,9 +92,13 @@ const props = defineProps({
|
|||||||
},
|
},
|
||||||
// 是否显示推理过程
|
// 是否显示推理过程
|
||||||
showRefs: {
|
showRefs: {
|
||||||
|
type: [Array, Boolean],
|
||||||
|
default: () => false
|
||||||
|
},
|
||||||
|
debugMode: {
|
||||||
type: Boolean,
|
type: Boolean,
|
||||||
default: false
|
default: false
|
||||||
}
|
},
|
||||||
});
|
});
|
||||||
|
|
||||||
const statusDefination = {
|
const statusDefination = {
|
||||||
|
|||||||
@ -4,19 +4,23 @@
|
|||||||
<!-- <span class="item btn" @click="likeThisResponse(msg)"><LikeOutlined /></span> -->
|
<!-- <span class="item btn" @click="likeThisResponse(msg)"><LikeOutlined /></span> -->
|
||||||
<!-- <span class="item btn" @click="dislikeThisResponse(msg)"><DislikeOutlined /></span> -->
|
<!-- <span class="item btn" @click="dislikeThisResponse(msg)"><DislikeOutlined /></span> -->
|
||||||
<span v-if="msg.meta?.server_model_name" class="item"><BulbOutlined /> {{ msg.meta.server_model_name }}</span>
|
<span v-if="msg.meta?.server_model_name" class="item"><BulbOutlined /> {{ msg.meta.server_model_name }}</span>
|
||||||
<span class="item btn" @click="copyText(msg.content)" title="复制"><CopyOutlined /></span>
|
|
||||||
<span class="item btn" @click="regenerateMessage()" title="重新生成"><ReloadOutlined /></span>
|
|
||||||
<span
|
<span
|
||||||
|
v-if="showKey('copy')"
|
||||||
|
class="item btn" @click="copyText(msg.content)" title="复制"><CopyOutlined /></span>
|
||||||
|
<span
|
||||||
|
v-if="showKey('regenerate')"
|
||||||
|
class="item btn" @click="regenerateMessage()" title="重新生成"><ReloadOutlined /></span>
|
||||||
|
<span
|
||||||
|
v-if="showKey('subGraph') && hasSubGraphData(msg)"
|
||||||
class="item btn"
|
class="item btn"
|
||||||
@click="openSubGraph(msg)"
|
@click="openSubGraph(msg)"
|
||||||
v-if="hasSubGraphData(msg)"
|
|
||||||
>
|
>
|
||||||
<DeploymentUnitOutlined /> 关系图
|
<DeploymentUnitOutlined /> 关系图
|
||||||
</span>
|
</span>
|
||||||
<span
|
<span
|
||||||
class="item btn"
|
class="item btn"
|
||||||
@click="showWebResult(msg)"
|
@click="showWebResult(msg)"
|
||||||
v-if="msg.refs?.web_search.results.length > 0"
|
v-if="showKey('webSearch') && msg.refs?.web_search.results.length > 0"
|
||||||
>
|
>
|
||||||
<GlobalOutlined /> 网页搜索 {{ msg.refs.web_search?.results.length }}
|
<GlobalOutlined /> 网页搜索 {{ msg.refs.web_search?.results.length }}
|
||||||
</span>
|
</span>
|
||||||
@ -120,6 +124,10 @@ import GraphContainer from './GraphContainer.vue' // 导入 GraphContainer 组
|
|||||||
const emit = defineEmits(['retry']);
|
const emit = defineEmits(['retry']);
|
||||||
const props = defineProps({
|
const props = defineProps({
|
||||||
message: Object,
|
message: Object,
|
||||||
|
showRefs: {
|
||||||
|
type: [Array, Boolean],
|
||||||
|
default: () => false
|
||||||
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
const msg = ref(props.message)
|
const msg = ref(props.message)
|
||||||
@ -127,6 +135,13 @@ const msg = ref(props.message)
|
|||||||
// 使用 useClipboard 实现复制功能
|
// 使用 useClipboard 实现复制功能
|
||||||
const { copy, isSupported } = useClipboard()
|
const { copy, isSupported } = useClipboard()
|
||||||
|
|
||||||
|
const showKey = (key) => {
|
||||||
|
if (props.showRefs === true) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return props.showRefs.includes(key)
|
||||||
|
}
|
||||||
|
|
||||||
// 定义 copy 方法
|
// 定义 copy 方法
|
||||||
const copyText = async (text) => {
|
const copyText = async (text) => {
|
||||||
if (isSupported) {
|
if (isSupported) {
|
||||||
|
|||||||
258
web/src/components/TokenManagerComponent.vue
Normal file
258
web/src/components/TokenManagerComponent.vue
Normal file
@ -0,0 +1,258 @@
|
|||||||
|
<template>
|
||||||
|
<div class="token-manager">
|
||||||
|
<div class="token-tools">
|
||||||
|
<a-button type="primary" size="small" @click="showAddTokenModal">
|
||||||
|
<PlusOutlined /> 创建 Token
|
||||||
|
</a-button>
|
||||||
|
</div>
|
||||||
|
<!-- 令牌列表 -->
|
||||||
|
<div class="token-list" v-if="tokens.length > 0">
|
||||||
|
<a-spin :spinning="loading">
|
||||||
|
<a-list size="small">
|
||||||
|
<a-list-item v-for="token in tokens" :key="token.id">
|
||||||
|
<div class="token-item">
|
||||||
|
<div class="token-info">
|
||||||
|
<div class="token-name">{{ token.name }}</div>
|
||||||
|
<div class="token-value">
|
||||||
|
<code>{{ token.token }}</code>
|
||||||
|
<a-button type="link" size="small" @click="copyToken(token.token)">
|
||||||
|
<CopyOutlined />
|
||||||
|
</a-button>
|
||||||
|
</div>
|
||||||
|
<div class="token-time">创建时间: {{ formatDate(token.created_at) }}</div>
|
||||||
|
</div>
|
||||||
|
<div class="token-actions">
|
||||||
|
<a-popconfirm
|
||||||
|
title="确定要删除这个令牌吗?"
|
||||||
|
ok-text="确定"
|
||||||
|
cancel-text="取消"
|
||||||
|
@confirm="deleteToken(token.id)"
|
||||||
|
>
|
||||||
|
<a-button type="text" danger size="small">
|
||||||
|
<DeleteOutlined />
|
||||||
|
</a-button>
|
||||||
|
</a-popconfirm>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</a-list-item>
|
||||||
|
</a-list>
|
||||||
|
</a-spin>
|
||||||
|
</div>
|
||||||
|
<a-empty v-else description="暂无访问令牌" :image="Empty.PRESENTED_IMAGE_SIMPLE" />
|
||||||
|
|
||||||
|
<!-- 添加令牌弹窗 -->
|
||||||
|
<a-modal
|
||||||
|
v-model:open="addTokenModalVisible"
|
||||||
|
title="添加访问令牌"
|
||||||
|
ok-text="创建"
|
||||||
|
cancel-text="取消"
|
||||||
|
@ok="createToken"
|
||||||
|
>
|
||||||
|
<a-form :model="newToken" layout="vertical">
|
||||||
|
<a-form-item label="令牌名称" name="name">
|
||||||
|
<a-input v-model:value="newToken.name" placeholder="请输入令牌名称" />
|
||||||
|
</a-form-item>
|
||||||
|
</a-form>
|
||||||
|
</a-modal>
|
||||||
|
</div>
|
||||||
|
</template>
|
||||||
|
|
||||||
|
<script setup>
|
||||||
|
import { ref, onMounted, watch } from 'vue';
|
||||||
|
import { message, Empty } from 'ant-design-vue';
|
||||||
|
import { PlusOutlined, DeleteOutlined, CopyOutlined } from '@ant-design/icons-vue';
|
||||||
|
|
||||||
|
const props = defineProps({
|
||||||
|
agentId: {
|
||||||
|
type: String,
|
||||||
|
required: true
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
// 状态
|
||||||
|
const tokens = ref([]);
|
||||||
|
const loading = ref(false);
|
||||||
|
const addTokenModalVisible = ref(false);
|
||||||
|
const newToken = ref({
|
||||||
|
name: ''
|
||||||
|
});
|
||||||
|
|
||||||
|
// 获取令牌列表
|
||||||
|
const fetchTokens = async () => {
|
||||||
|
loading.value = true;
|
||||||
|
try {
|
||||||
|
const response = await fetch(`/api/admin/tokens?agent_id=${props.agentId}`);
|
||||||
|
if (response.ok) {
|
||||||
|
const data = await response.json();
|
||||||
|
tokens.value = data;
|
||||||
|
} else {
|
||||||
|
message.error('获取令牌列表失败');
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
console.error('获取令牌列表出错:', error);
|
||||||
|
message.error('获取令牌列表出错');
|
||||||
|
} finally {
|
||||||
|
loading.value = false;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// 创建新令牌
|
||||||
|
const createToken = async () => {
|
||||||
|
if (!newToken.value.name.trim()) {
|
||||||
|
message.warning('请输入令牌名称');
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
try {
|
||||||
|
const response = await fetch('/api/admin/tokens', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: {
|
||||||
|
'Content-Type': 'application/json'
|
||||||
|
},
|
||||||
|
body: JSON.stringify({
|
||||||
|
agent_id: props.agentId,
|
||||||
|
name: newToken.value.name
|
||||||
|
})
|
||||||
|
});
|
||||||
|
|
||||||
|
if (response.ok) {
|
||||||
|
const data = await response.json();
|
||||||
|
tokens.value.push(data);
|
||||||
|
message.success('令牌创建成功');
|
||||||
|
addTokenModalVisible.value = false;
|
||||||
|
newToken.value.name = '';
|
||||||
|
} else {
|
||||||
|
message.error('创建令牌失败');
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
console.error('创建令牌出错:', error);
|
||||||
|
message.error('创建令牌出错');
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// 删除令牌
|
||||||
|
const deleteToken = async (tokenId) => {
|
||||||
|
try {
|
||||||
|
const response = await fetch(`/api/admin/tokens/${tokenId}`, {
|
||||||
|
method: 'DELETE'
|
||||||
|
});
|
||||||
|
|
||||||
|
if (response.ok) {
|
||||||
|
tokens.value = tokens.value.filter(token => token.id !== tokenId);
|
||||||
|
message.success('令牌已删除');
|
||||||
|
} else {
|
||||||
|
message.error('删除令牌失败');
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
console.error('删除令牌出错:', error);
|
||||||
|
message.error('删除令牌出错');
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// 复制令牌到剪贴板
|
||||||
|
const copyToken = (token) => {
|
||||||
|
navigator.clipboard.writeText(token).then(() => {
|
||||||
|
message.success('令牌已复制到剪贴板');
|
||||||
|
});
|
||||||
|
};
|
||||||
|
|
||||||
|
// 显示添加令牌弹窗
|
||||||
|
const showAddTokenModal = () => {
|
||||||
|
newToken.value.name = '';
|
||||||
|
addTokenModalVisible.value = true;
|
||||||
|
};
|
||||||
|
|
||||||
|
// 格式化日期
|
||||||
|
const formatDate = (dateString) => {
|
||||||
|
if (!dateString) return '';
|
||||||
|
const date = new Date(dateString);
|
||||||
|
return date.toLocaleString();
|
||||||
|
};
|
||||||
|
|
||||||
|
// 监听agentId变化
|
||||||
|
watch(() => props.agentId, (newAgentId) => {
|
||||||
|
if (newAgentId) {
|
||||||
|
fetchTokens();
|
||||||
|
} else {
|
||||||
|
tokens.value = [];
|
||||||
|
}
|
||||||
|
}, { immediate: true });
|
||||||
|
|
||||||
|
// 组件挂载时获取令牌列表
|
||||||
|
onMounted(() => {
|
||||||
|
if (props.agentId) {
|
||||||
|
fetchTokens();
|
||||||
|
}
|
||||||
|
});
|
||||||
|
</script>
|
||||||
|
|
||||||
|
<style lang="less" scoped>
|
||||||
|
.token-manager {
|
||||||
|
margin-top: 1rem;
|
||||||
|
// padding: 0 0.5rem;
|
||||||
|
}
|
||||||
|
|
||||||
|
.manager-title {
|
||||||
|
font-size: 1rem;
|
||||||
|
margin-bottom: 1rem;
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-tools {
|
||||||
|
margin-bottom: 1rem;
|
||||||
|
display: flex;
|
||||||
|
justify-content: flex-end;
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-list {
|
||||||
|
max-height: calc(100vh - 400px);
|
||||||
|
overflow-y: auto;
|
||||||
|
li.ant-list-item {
|
||||||
|
background-color: var(--gray-100);
|
||||||
|
border-radius: 0.5rem;
|
||||||
|
padding: 0.5rem;
|
||||||
|
margin-bottom: 0.5rem;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-item {
|
||||||
|
display: flex;
|
||||||
|
justify-content: space-between;
|
||||||
|
width: 100%;
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-info {
|
||||||
|
flex: 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-name {
|
||||||
|
font-weight: 500;
|
||||||
|
margin-bottom: 0.25rem;
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-value {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
gap: 0.5rem;
|
||||||
|
font-size: 0.8rem;
|
||||||
|
background-color: var(--main-light-4);
|
||||||
|
padding: 0.25rem 0.5rem;
|
||||||
|
border-radius: 4px;
|
||||||
|
margin-bottom: 0.25rem;
|
||||||
|
overflow-x: auto;
|
||||||
|
user-select: all;
|
||||||
|
|
||||||
|
code {
|
||||||
|
flex: 1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-time {
|
||||||
|
font-size: 0.75rem;
|
||||||
|
color: var(--gray-500);
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-actions {
|
||||||
|
display: flex;
|
||||||
|
align-items: center;
|
||||||
|
}
|
||||||
|
</style>
|
||||||
@ -31,29 +31,29 @@ const router = createRouter({
|
|||||||
}
|
}
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
// {
|
||||||
path: '/agent',
|
// path: '/agent',
|
||||||
name: 'agent',
|
// name: 'agent',
|
||||||
component: AppLayout,
|
// component: AppLayout,
|
||||||
children: [
|
// children: [
|
||||||
{
|
// {
|
||||||
path: '',
|
// path: '',
|
||||||
name: 'AgentMain',
|
// name: 'AgentMain',
|
||||||
component: () => import('../views/AgentView.vue'),
|
// component: () => import('../views/AgentView.vue'),
|
||||||
meta: { keepAlive: true }
|
// meta: { keepAlive: true }
|
||||||
},
|
// },
|
||||||
{
|
// {
|
||||||
path: ':agent_id',
|
// path: ':agent_id',
|
||||||
name: 'AgentSinglePage',
|
// name: 'AgentSinglePage',
|
||||||
component: () => import('../components/AgentSingleViewComponent.vue'),
|
// component: () => import('../components/AgentSingleViewComponent.vue'),
|
||||||
meta: { keepAlive: false }
|
// meta: { keepAlive: false }
|
||||||
}
|
// }
|
||||||
]
|
// ]
|
||||||
},
|
// },
|
||||||
{
|
{
|
||||||
path: '/agent/:agent_id',
|
path: '/agent/:agent_id',
|
||||||
name: 'AgentSinglePage',
|
name: 'AgentSinglePage',
|
||||||
component: () => import('../views/AgentView.vue'),
|
component: () => import('../views/AgentSingleView.vue'),
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
path: '/graph',
|
path: '/graph',
|
||||||
|
|||||||
163
web/src/views/AgentSingleView.vue
Normal file
163
web/src/views/AgentSingleView.vue
Normal file
@ -0,0 +1,163 @@
|
|||||||
|
<template>
|
||||||
|
<div class="agent-single-view">
|
||||||
|
<!-- Token验证弹窗 -->
|
||||||
|
<a-modal
|
||||||
|
v-model:open="tokenModalVisible"
|
||||||
|
title="访问验证"
|
||||||
|
:closable="false"
|
||||||
|
:maskClosable="false"
|
||||||
|
:keyboard="false"
|
||||||
|
:footer="null"
|
||||||
|
width="500px"
|
||||||
|
>
|
||||||
|
<div class="token-verify-form">
|
||||||
|
<p>需要输入访问令牌才能使用该智能体</p>
|
||||||
|
<a-input-password
|
||||||
|
v-model:value="tokenInput"
|
||||||
|
placeholder="请输入访问令牌"
|
||||||
|
@pressEnter="verifyToken"
|
||||||
|
/>
|
||||||
|
<div class="error-message" v-if="errorMessage">{{ errorMessage }}</div>
|
||||||
|
<div class="token-actions">
|
||||||
|
<a-button type="primary" :loading="verifying" @click="verifyToken">验证</a-button>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
</a-modal>
|
||||||
|
|
||||||
|
<!-- 智能体聊天界面 -->
|
||||||
|
<AgentChatComponent v-if="isVerified" :agent-id="agentId" />
|
||||||
|
</div>
|
||||||
|
</template>
|
||||||
|
|
||||||
|
<script setup>
|
||||||
|
import { ref, computed, onMounted } from 'vue';
|
||||||
|
import { useRoute, useRouter } from 'vue-router';
|
||||||
|
import AgentChatComponent from '@/components/AgentChatComponent.vue';
|
||||||
|
import { message } from 'ant-design-vue';
|
||||||
|
|
||||||
|
const route = useRoute();
|
||||||
|
const router = useRouter();
|
||||||
|
const agentId = computed(() => route.params.agent_id);
|
||||||
|
|
||||||
|
// Token验证相关状态
|
||||||
|
const tokenModalVisible = ref(false);
|
||||||
|
const tokenInput = ref('');
|
||||||
|
const isVerified = ref(false);
|
||||||
|
const verifying = ref(false);
|
||||||
|
const errorMessage = ref('');
|
||||||
|
|
||||||
|
// 验证Token
|
||||||
|
const verifyToken = async () => {
|
||||||
|
if (!tokenInput.value) {
|
||||||
|
errorMessage.value = '请输入访问令牌';
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
verifying.value = true;
|
||||||
|
errorMessage.value = '';
|
||||||
|
|
||||||
|
try {
|
||||||
|
const response = await fetch('/api/admin/verify_token', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: {
|
||||||
|
'Content-Type': 'application/json'
|
||||||
|
},
|
||||||
|
body: JSON.stringify({
|
||||||
|
agent_id: agentId.value,
|
||||||
|
token: tokenInput.value
|
||||||
|
})
|
||||||
|
});
|
||||||
|
|
||||||
|
if (response.ok) {
|
||||||
|
// 验证成功
|
||||||
|
isVerified.value = true;
|
||||||
|
tokenModalVisible.value = false;
|
||||||
|
|
||||||
|
// 保存令牌到localStorage
|
||||||
|
localStorage.setItem(`agent-token-${agentId.value}`, tokenInput.value);
|
||||||
|
} else {
|
||||||
|
// 验证失败
|
||||||
|
const data = await response.json();
|
||||||
|
errorMessage.value = data.detail || '令牌验证失败';
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
console.error('验证令牌出错:', error);
|
||||||
|
errorMessage.value = '验证令牌时发生错误';
|
||||||
|
} finally {
|
||||||
|
verifying.value = false;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// 检查是否已经验证
|
||||||
|
const checkVerification = async () => {
|
||||||
|
const savedToken = localStorage.getItem(`agent-token-${agentId.value}`);
|
||||||
|
|
||||||
|
if (savedToken) {
|
||||||
|
// 即使有保存的令牌,也要重新验证其有效性
|
||||||
|
verifying.value = true;
|
||||||
|
|
||||||
|
try {
|
||||||
|
const response = await fetch('/api/admin/verify_token', {
|
||||||
|
method: 'POST',
|
||||||
|
headers: {
|
||||||
|
'Content-Type': 'application/json'
|
||||||
|
},
|
||||||
|
body: JSON.stringify({
|
||||||
|
agent_id: agentId.value,
|
||||||
|
token: savedToken
|
||||||
|
})
|
||||||
|
});
|
||||||
|
|
||||||
|
if (response.ok) {
|
||||||
|
// 验证成功
|
||||||
|
isVerified.value = true;
|
||||||
|
tokenInput.value = savedToken; // 保存已验证的令牌到输入框
|
||||||
|
} else {
|
||||||
|
// 令牌无效,清除本地存储并显示输入框
|
||||||
|
localStorage.removeItem(`agent-token-${agentId.value}`);
|
||||||
|
tokenModalVisible.value = true;
|
||||||
|
}
|
||||||
|
} catch (error) {
|
||||||
|
console.error('验证令牌出错:', error);
|
||||||
|
tokenModalVisible.value = true;
|
||||||
|
} finally {
|
||||||
|
verifying.value = false;
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
// 没有保存的令牌,显示输入框
|
||||||
|
tokenModalVisible.value = true;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// 组件挂载时检查验证状态
|
||||||
|
onMounted(() => {
|
||||||
|
checkVerification();
|
||||||
|
});
|
||||||
|
</script>
|
||||||
|
|
||||||
|
<style lang="less" scoped>
|
||||||
|
.agent-single-view {
|
||||||
|
width: 100%;
|
||||||
|
height: 100vh;
|
||||||
|
overflow: hidden;
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-verify-form {
|
||||||
|
display: flex;
|
||||||
|
flex-direction: column;
|
||||||
|
gap: 1rem;
|
||||||
|
|
||||||
|
.error-message {
|
||||||
|
color: #ff4d4f;
|
||||||
|
font-size: 0.85rem;
|
||||||
|
}
|
||||||
|
|
||||||
|
.token-actions {
|
||||||
|
display: flex;
|
||||||
|
justify-content: space-between;
|
||||||
|
margin-top: 0.5rem;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
</style>
|
||||||
|
|
||||||
|
|
||||||
File diff suppressed because it is too large
Load Diff
Loading…
Reference in New Issue
Block a user