feat: initial release - Hunzi agent framework v0.1.0
This commit is contained in:
@@ -0,0 +1,234 @@
|
||||
"""Agent 管理 API"""
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.database import get_session
|
||||
from app.models import Agent, AgentTool
|
||||
from app.config import get_config
|
||||
from app.llm.client import LLMClient
|
||||
from app.agent.engine import Agent as AgentEngine
|
||||
from app.agent.tools import global_registry
|
||||
from app.agent.memory import ConversationMemory
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
# --- Pydantic Models ---
|
||||
|
||||
class AgentCreate(BaseModel):
|
||||
name: str
|
||||
description: str = ""
|
||||
system_prompt: str = "You are a helpful assistant."
|
||||
llm_provider: str = "default"
|
||||
temperature: float = 0.7
|
||||
max_tokens: int = 4096
|
||||
tool_ids: List[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class AgentUpdate(BaseModel):
|
||||
name: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
system_prompt: Optional[str] = None
|
||||
llm_provider: Optional[str] = None
|
||||
temperature: Optional[float] = None
|
||||
max_tokens: Optional[int] = None
|
||||
tool_ids: Optional[List[str]] = None
|
||||
|
||||
|
||||
class ChatRequest(BaseModel):
|
||||
message: str
|
||||
images: List[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class AgentResponse(BaseModel):
|
||||
id: str
|
||||
name: str
|
||||
description: str
|
||||
system_prompt: str
|
||||
llm_provider: str
|
||||
temperature: float
|
||||
max_tokens: int
|
||||
tools: List[str]
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
def _agent_to_response(agent: Agent) -> AgentResponse:
|
||||
return AgentResponse(
|
||||
id=agent.id,
|
||||
name=agent.name,
|
||||
description=agent.description or "",
|
||||
system_prompt=agent.system_prompt or "",
|
||||
llm_provider=agent.llm_provider or "default",
|
||||
temperature=agent.temperature,
|
||||
max_tokens=agent.max_tokens,
|
||||
tools=[t.tool_name for t in agent.tools if t.enabled],
|
||||
created_at=agent.created_at,
|
||||
updated_at=agent.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _build_agent_engine(agent: Agent) -> AgentEngine:
|
||||
"""从数据库 Agent 构建运行时的 Agent 引擎"""
|
||||
config = get_config()
|
||||
provider_key = agent.llm_provider or "default"
|
||||
provider_cfg = config.llm_providers.get(provider_key) or config.default_llm
|
||||
|
||||
llm = LLMClient(provider_cfg)
|
||||
tools = global_registry
|
||||
memory = ConversationMemory(window_size=50)
|
||||
|
||||
return AgentEngine(
|
||||
name=agent.name,
|
||||
llm_client=llm,
|
||||
tools=tools,
|
||||
memory=memory,
|
||||
system_prompt=agent.system_prompt or "You are a helpful assistant.",
|
||||
temperature=agent.temperature,
|
||||
max_tokens=agent.max_tokens,
|
||||
)
|
||||
|
||||
|
||||
# --- Routes ---
|
||||
|
||||
@router.get("")
|
||||
async def list_agents() -> List[AgentResponse]:
|
||||
session = get_session()
|
||||
try:
|
||||
agents = session.query(Agent).order_by(Agent.created_at.desc()).all()
|
||||
return [_agent_to_response(a) for a in agents]
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
@router.post("", response_model=AgentResponse)
|
||||
async def create_agent(data: AgentCreate) -> AgentResponse:
|
||||
session = get_session()
|
||||
try:
|
||||
agent = Agent(
|
||||
id=str(uuid.uuid4()),
|
||||
name=data.name,
|
||||
description=data.description,
|
||||
system_prompt=data.system_prompt,
|
||||
llm_provider=data.llm_provider,
|
||||
temperature=data.temperature,
|
||||
max_tokens=data.max_tokens,
|
||||
)
|
||||
for tool_name in data.tool_ids:
|
||||
agent.tools.append(AgentTool(tool_name=tool_name, enabled=True))
|
||||
session.add(agent)
|
||||
session.commit()
|
||||
session.refresh(agent)
|
||||
return _agent_to_response(agent)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
@router.get("/{agent_id}", response_model=AgentResponse)
|
||||
async def get_agent(agent_id: str) -> AgentResponse:
|
||||
session = get_session()
|
||||
try:
|
||||
agent = session.query(Agent).filter(Agent.id == agent_id).first()
|
||||
if not agent:
|
||||
raise HTTPException(404, "Agent not found")
|
||||
return _agent_to_response(agent)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
@router.put("/{agent_id}", response_model=AgentResponse)
|
||||
async def update_agent(agent_id: str, data: AgentUpdate) -> AgentResponse:
|
||||
session = get_session()
|
||||
try:
|
||||
agent = session.query(Agent).filter(Agent.id == agent_id).first()
|
||||
if not agent:
|
||||
raise HTTPException(404, "Agent not found")
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
for key, value in update_data.items():
|
||||
if value is not None and key != "tool_ids":
|
||||
setattr(agent, key, value)
|
||||
|
||||
if "tool_ids" in update_data and update_data["tool_ids"] is not None:
|
||||
agent.tools = [
|
||||
AgentTool(tool_name=t, enabled=True) for t in update_data["tool_ids"]
|
||||
]
|
||||
|
||||
session.commit()
|
||||
session.refresh(agent)
|
||||
return _agent_to_response(agent)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
@router.delete("/{agent_id}")
|
||||
async def delete_agent(agent_id: str) -> dict:
|
||||
session = get_session()
|
||||
try:
|
||||
agent = session.query(Agent).filter(Agent.id == agent_id).first()
|
||||
if not agent:
|
||||
raise HTTPException(404, "Agent not found")
|
||||
session.delete(agent)
|
||||
session.commit()
|
||||
return {"deleted": agent_id}
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
@router.post("/{agent_id}/chat")
|
||||
async def chat_with_agent(agent_id: str, data: ChatRequest) -> dict:
|
||||
session = get_session()
|
||||
try:
|
||||
agent = session.query(Agent).filter(Agent.id == agent_id).first()
|
||||
if not agent:
|
||||
raise HTTPException(404, "Agent not found")
|
||||
|
||||
engine = _build_agent_engine(agent)
|
||||
|
||||
messages = [{"role": "user", "content": data.message}]
|
||||
if data.images:
|
||||
for img in data.images:
|
||||
messages[0]["content"] = [
|
||||
{"type": "text", "text": data.message},
|
||||
]
|
||||
if img.startswith(("http://", "https://")):
|
||||
messages[0]["content"].append({
|
||||
"type": "image_url",
|
||||
"image_url": {"url": img},
|
||||
})
|
||||
else:
|
||||
messages[0]["content"].append({
|
||||
"type": "image_url",
|
||||
"image_url": {"url": f"data:image/png;base64,{img}"},
|
||||
})
|
||||
|
||||
response = engine.run(messages)
|
||||
return {"response": response, "agent_id": agent_id}
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
@router.post("/{agent_id}/chat/stream")
|
||||
async def chat_stream(agent_id: str, data: ChatRequest) -> StreamingResponse:
|
||||
session = get_session()
|
||||
try:
|
||||
agent = session.query(Agent).filter(Agent.id == agent_id).first()
|
||||
if not agent:
|
||||
raise HTTPException(404, "Agent not found")
|
||||
|
||||
engine = _build_agent_engine(agent)
|
||||
|
||||
async def event_stream():
|
||||
messages = [{"role": "user", "content": data.message}]
|
||||
for chunk in engine.run_stream(messages):
|
||||
yield f"data: {chunk}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
finally:
|
||||
session.close()
|
||||
Reference in New Issue
Block a user