619 lines
23 KiB
Python
619 lines
23 KiB
Python
"""
|
|
LiteLLM Agent 核心模块
|
|
|
|
基于LiteLLM框架的Agent实现,支持A2A协议
|
|
"""
|
|
import asyncio
|
|
import json
|
|
import uuid
|
|
from typing import AsyncGenerator, Optional, Dict, Any, List, Union
|
|
from dataclasses import dataclass, field
|
|
from datetime import datetime
|
|
|
|
import httpx
|
|
import structlog
|
|
|
|
from config import LiteLLMConfig, AgentConfig, get_config
|
|
|
|
# 配置日志
|
|
logger = structlog.get_logger()
|
|
|
|
|
|
@dataclass
|
|
class Message:
|
|
"""消息数据结构"""
|
|
role: str # user, assistant, system
|
|
content: str
|
|
message_id: str = field(default_factory=lambda: uuid.uuid4().hex)
|
|
timestamp: datetime = field(default_factory=datetime.now)
|
|
metadata: Dict[str, Any] = field(default_factory=dict)
|
|
|
|
|
|
@dataclass
|
|
class Conversation:
|
|
"""对话上下文"""
|
|
conversation_id: str = field(default_factory=lambda: uuid.uuid4().hex)
|
|
messages: List[Message] = field(default_factory=list)
|
|
created_at: datetime = field(default_factory=datetime.now)
|
|
|
|
def add_message(self, role: str, content: str) -> Message:
|
|
"""添加消息到对话"""
|
|
msg = Message(role=role, content=content)
|
|
self.messages.append(msg)
|
|
return msg
|
|
|
|
def to_openai_format(self) -> List[Dict[str, str]]:
|
|
"""转换为OpenAI格式的消息列表"""
|
|
return [{"role": m.role, "content": m.content} for m in self.messages]
|
|
|
|
|
|
@dataclass
|
|
class ModelResult:
|
|
"""Normalized model response metadata for Runtime accounting."""
|
|
|
|
content: str
|
|
usage: Dict[str, int] = field(default_factory=dict)
|
|
request_id: Optional[str] = None
|
|
response_id: Optional[str] = None
|
|
model: Optional[str] = None
|
|
api_format: str = "openai_chat"
|
|
endpoint: Optional[str] = None
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
return {
|
|
"content": self.content,
|
|
"usage": self.usage,
|
|
"request_id": self.request_id,
|
|
"response_id": self.response_id,
|
|
"model": self.model,
|
|
"api_format": self.api_format,
|
|
"endpoint": self.endpoint,
|
|
}
|
|
|
|
|
|
class ModelRequestError(RuntimeError):
|
|
"""Model gateway error carrying request metadata for Runtime logs."""
|
|
|
|
def __init__(
|
|
self,
|
|
message: str,
|
|
*,
|
|
request_id: Optional[str] = None,
|
|
status_code: Optional[int] = None,
|
|
response_text: Optional[str] = None,
|
|
):
|
|
super().__init__(message)
|
|
self.request_id = request_id
|
|
self.status_code = status_code
|
|
self.response_text = response_text
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
return {
|
|
"request_id": self.request_id,
|
|
"status_code": self.status_code,
|
|
"response_text": self.response_text,
|
|
}
|
|
|
|
|
|
class LiteLLMAgent:
|
|
"""
|
|
基于LiteLLM的Agent实现
|
|
|
|
支持功能:
|
|
- 多轮对话
|
|
- 流式响应
|
|
- 工具调用
|
|
- A2A协议兼容
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
api_key: Optional[str] = None,
|
|
model: Optional[str] = None,
|
|
litellm_config: Optional[LiteLLMConfig] = None,
|
|
agent_config: Optional[AgentConfig] = None
|
|
):
|
|
"""
|
|
初始化Agent
|
|
|
|
Args:
|
|
api_key: LiteLLM API密钥(可选,优先使用,否则从环境变量获取)
|
|
model: 模型名称(可选,优先使用,否则从环境变量获取)
|
|
litellm_config: LiteLLM配置对象
|
|
agent_config: Agent配置对象
|
|
"""
|
|
if litellm_config:
|
|
self.llm_config = litellm_config
|
|
else:
|
|
self.llm_config = LiteLLMConfig(api_key=api_key, model=model)
|
|
|
|
if agent_config:
|
|
self.agent_config = agent_config
|
|
else:
|
|
self.agent_config = AgentConfig()
|
|
|
|
# 验证配置
|
|
self.llm_config.validate()
|
|
|
|
# HTTP客户端
|
|
self._client: Optional[httpx.AsyncClient] = None
|
|
|
|
# 对话管理
|
|
self.conversations: Dict[str, Conversation] = {}
|
|
|
|
# 工具注册
|
|
self.tools: Dict[str, callable] = {}
|
|
|
|
logger.info(
|
|
"Agent初始化完成",
|
|
agent_name=self.agent_config.name,
|
|
model=self.llm_config.model,
|
|
base_url=self.llm_config.base_url
|
|
)
|
|
|
|
async def _get_client(self) -> httpx.AsyncClient:
|
|
"""获取或创建HTTP客户端"""
|
|
if self._client is None or self._client.is_closed:
|
|
self._client = httpx.AsyncClient(
|
|
timeout=httpx.Timeout(self.llm_config.timeout),
|
|
headers={
|
|
"Authorization": f"Bearer {self.llm_config.api_key}",
|
|
"x-api-key": self.llm_config.api_key,
|
|
"anthropic-version": "2023-06-01",
|
|
"Content-Type": "application/json"
|
|
}
|
|
)
|
|
return self._client
|
|
|
|
async def close(self):
|
|
"""关闭资源"""
|
|
if self._client and not self._client.is_closed:
|
|
await self._client.aclose()
|
|
|
|
def register_tool(self, name: str, func: callable, description: str = ""):
|
|
"""注册工具函数"""
|
|
self.tools[name] = {
|
|
"function": func,
|
|
"description": description
|
|
}
|
|
logger.info(f"注册工具: {name}")
|
|
|
|
def get_or_create_conversation(self, conversation_id: Optional[str] = None) -> Conversation:
|
|
"""获取或创建对话"""
|
|
if conversation_id and conversation_id in self.conversations:
|
|
return self.conversations[conversation_id]
|
|
|
|
conv = Conversation(conversation_id=conversation_id or uuid.uuid4().hex)
|
|
# 添加系统提示
|
|
conv.add_message("system", self.agent_config.system_prompt)
|
|
self.conversations[conv.conversation_id] = conv
|
|
return conv
|
|
|
|
async def chat(
|
|
self,
|
|
message: str,
|
|
conversation_id: Optional[str] = None,
|
|
stream: bool = False
|
|
) -> Union[str, AsyncGenerator[str, None]]:
|
|
"""
|
|
发送消息并获取回复
|
|
|
|
Args:
|
|
message: 用户消息
|
|
conversation_id: 对话ID(用于多轮对话)
|
|
stream: 是否流式响应
|
|
|
|
Returns:
|
|
如果stream=False,返回完整回复字符串
|
|
如果stream=True,返回异步生成器
|
|
"""
|
|
# 获取对话上下文
|
|
conversation = self.get_or_create_conversation(conversation_id)
|
|
conversation.add_message("user", message)
|
|
|
|
if stream:
|
|
if self.llm_config.api_format == "anthropic_messages":
|
|
return self._stream_anthropic_messages_text(conversation)
|
|
return self._stream_chat(conversation)
|
|
result = await self.chat_result_for_conversation(conversation)
|
|
return result.content
|
|
|
|
async def chat_result(
|
|
self,
|
|
message: str,
|
|
conversation_id: Optional[str] = None,
|
|
) -> Dict[str, Any]:
|
|
"""Return assistant text plus usage and NewAPI request metadata."""
|
|
conversation = self.get_or_create_conversation(conversation_id)
|
|
conversation.add_message("user", message)
|
|
return (await self.chat_result_for_conversation(conversation)).to_dict()
|
|
|
|
async def chat_result_for_conversation(self, conversation: Conversation) -> ModelResult:
|
|
"""Dispatch to the configured model API format."""
|
|
if self.llm_config.api_format == "anthropic_messages":
|
|
if self.llm_config.use_stream:
|
|
return await self._anthropic_messages_stream(conversation)
|
|
return await self._anthropic_messages(conversation)
|
|
if self.llm_config.use_stream:
|
|
return await self._openai_chat_stream_result(conversation)
|
|
return await self._simple_chat(conversation)
|
|
|
|
def _request_id_from_response(self, response: httpx.Response, body: Optional[Dict[str, Any]] = None) -> Optional[str]:
|
|
"""Extract NewAPI/OpenAI/Anthropic request ID from headers or body."""
|
|
for name in (
|
|
"x-request-id",
|
|
"request-id",
|
|
"x-newapi-request-id",
|
|
"x-litellm-request-id",
|
|
"anthropic-request-id",
|
|
):
|
|
value = response.headers.get(name)
|
|
if value:
|
|
return value
|
|
if body:
|
|
return body.get("request_id")
|
|
return None
|
|
|
|
def _normalize_usage(self, usage: Optional[Dict[str, Any]]) -> Dict[str, int]:
|
|
usage = usage or {}
|
|
prompt_tokens = int(usage.get("prompt_tokens") or usage.get("input_tokens") or 0)
|
|
completion_tokens = int(usage.get("completion_tokens") or usage.get("output_tokens") or 0)
|
|
if "input_tokens" in usage or "output_tokens" in usage:
|
|
total_tokens = prompt_tokens + completion_tokens
|
|
else:
|
|
total_tokens = int(usage.get("total_tokens") or prompt_tokens + completion_tokens)
|
|
return {
|
|
"prompt_tokens": prompt_tokens,
|
|
"completion_tokens": completion_tokens,
|
|
"total_tokens": total_tokens,
|
|
}
|
|
|
|
def _raise_gateway_error(self, exc: httpx.HTTPStatusError, body: Optional[Dict[str, Any]] = None) -> None:
|
|
request_id = self._request_id_from_response(exc.response, body)
|
|
raise ModelRequestError(
|
|
f"Server error '{exc.response.status_code} {exc.response.reason_phrase}' for url '{exc.request.url}'",
|
|
request_id=request_id,
|
|
status_code=exc.response.status_code,
|
|
response_text=exc.response.text[:2000],
|
|
) from exc
|
|
|
|
async def _simple_chat(self, conversation: Conversation) -> ModelResult:
|
|
"""非流式 OpenAI chat completions 对话"""
|
|
client = await self._get_client()
|
|
|
|
request_body = {
|
|
"model": self.llm_config.model,
|
|
"messages": conversation.to_openai_format(),
|
|
"temperature": self.llm_config.temperature,
|
|
"max_tokens": self.llm_config.max_tokens
|
|
}
|
|
|
|
logger.debug("发送请求", endpoint=self.llm_config.chat_endpoint)
|
|
|
|
try:
|
|
response = await client.post(
|
|
self.llm_config.chat_endpoint,
|
|
json=request_body
|
|
)
|
|
try:
|
|
response.raise_for_status()
|
|
except httpx.HTTPStatusError as exc:
|
|
self._raise_gateway_error(exc)
|
|
|
|
result = response.json()
|
|
assistant_message = result["choices"][0]["message"]["content"]
|
|
|
|
# 保存助手回复到对话
|
|
conversation.add_message("assistant", assistant_message)
|
|
|
|
request_id = self._request_id_from_response(response, result)
|
|
logger.info("收到回复", length=len(assistant_message), request_id=request_id)
|
|
return ModelResult(
|
|
content=assistant_message,
|
|
usage=self._normalize_usage(result.get("usage")),
|
|
request_id=request_id,
|
|
response_id=result.get("id"),
|
|
model=result.get("model") or self.llm_config.model,
|
|
api_format="openai_chat",
|
|
endpoint=self.llm_config.chat_endpoint,
|
|
)
|
|
|
|
except httpx.HTTPStatusError as e:
|
|
logger.error("HTTP错误", status_code=e.response.status_code, detail=e.response.text)
|
|
raise
|
|
except Exception as e:
|
|
logger.error("请求失败", error=str(e))
|
|
raise
|
|
|
|
async def _openai_chat_stream_result(self, conversation: Conversation) -> ModelResult:
|
|
"""OpenAI chat completions stream=true, aggregated into one Runtime artifact."""
|
|
client = await self._get_client()
|
|
request_body = {
|
|
"model": self.llm_config.model,
|
|
"messages": conversation.to_openai_format(),
|
|
"temperature": self.llm_config.temperature,
|
|
"max_tokens": self.llm_config.max_tokens,
|
|
"stream": True,
|
|
"stream_options": {"include_usage": True},
|
|
}
|
|
|
|
full_response = ""
|
|
usage: Dict[str, int] = {}
|
|
response_id: Optional[str] = None
|
|
request_id: Optional[str] = None
|
|
try:
|
|
async with client.stream("POST", self.llm_config.chat_endpoint, json=request_body) as response:
|
|
try:
|
|
response.raise_for_status()
|
|
except httpx.HTTPStatusError as exc:
|
|
self._raise_gateway_error(exc)
|
|
request_id = self._request_id_from_response(response)
|
|
async for line in response.aiter_lines():
|
|
if not line.startswith("data: "):
|
|
continue
|
|
data = line[6:]
|
|
if data == "[DONE]":
|
|
break
|
|
try:
|
|
chunk = json.loads(data)
|
|
except json.JSONDecodeError:
|
|
continue
|
|
response_id = response_id or chunk.get("id")
|
|
request_id = request_id or chunk.get("request_id")
|
|
if chunk.get("usage"):
|
|
usage = self._normalize_usage(chunk.get("usage"))
|
|
choices = chunk.get("choices") or []
|
|
if not choices:
|
|
continue
|
|
delta = choices[0].get("delta", {})
|
|
content = delta.get("content", "")
|
|
if content:
|
|
full_response += content
|
|
|
|
conversation.add_message("assistant", full_response)
|
|
logger.info("收到流式回复", length=len(full_response), request_id=request_id)
|
|
return ModelResult(
|
|
content=full_response,
|
|
usage=usage,
|
|
request_id=request_id or response_id,
|
|
response_id=response_id,
|
|
model=self.llm_config.model,
|
|
api_format="openai_chat",
|
|
endpoint=self.llm_config.chat_endpoint,
|
|
)
|
|
except Exception as e:
|
|
logger.error("流式请求失败", error=str(e))
|
|
raise
|
|
|
|
def _anthropic_payload(self, conversation: Conversation, *, stream: bool = False) -> Dict[str, Any]:
|
|
system_parts: List[str] = []
|
|
messages: List[Dict[str, str]] = []
|
|
for message in conversation.messages:
|
|
if message.role == "system":
|
|
system_parts.append(message.content)
|
|
else:
|
|
role = "assistant" if message.role == "assistant" else "user"
|
|
messages.append({"role": role, "content": message.content})
|
|
payload: Dict[str, Any] = {
|
|
"model": self.llm_config.model,
|
|
"messages": messages,
|
|
"max_tokens": self.llm_config.max_tokens,
|
|
"stream": stream,
|
|
}
|
|
if system_parts:
|
|
payload["system"] = "\n\n".join(system_parts)
|
|
return payload
|
|
|
|
def _anthropic_text(self, body: Dict[str, Any]) -> str:
|
|
content = body.get("content") or []
|
|
texts = [
|
|
part.get("text", "")
|
|
for part in content
|
|
if isinstance(part, dict) and part.get("type") == "text"
|
|
]
|
|
return "".join(texts)
|
|
|
|
async def _anthropic_messages(self, conversation: Conversation) -> ModelResult:
|
|
"""Anthropic Messages-compatible call for Claude models."""
|
|
client = await self._get_client()
|
|
try:
|
|
response = await client.post(self.llm_config.messages_endpoint, json=self._anthropic_payload(conversation))
|
|
try:
|
|
response.raise_for_status()
|
|
except httpx.HTTPStatusError as exc:
|
|
self._raise_gateway_error(exc)
|
|
result = response.json()
|
|
assistant_message = self._anthropic_text(result)
|
|
conversation.add_message("assistant", assistant_message)
|
|
request_id = self._request_id_from_response(response, result)
|
|
logger.info("收到 Claude Messages 回复", length=len(assistant_message), request_id=request_id)
|
|
return ModelResult(
|
|
content=assistant_message,
|
|
usage=self._normalize_usage(result.get("usage")),
|
|
request_id=request_id,
|
|
response_id=result.get("id"),
|
|
model=result.get("model") or self.llm_config.model,
|
|
api_format="anthropic_messages",
|
|
endpoint=self.llm_config.messages_endpoint,
|
|
)
|
|
except Exception as e:
|
|
logger.error("Claude Messages 请求失败", error=str(e))
|
|
raise
|
|
|
|
async def _anthropic_messages_stream(self, conversation: Conversation) -> ModelResult:
|
|
"""Anthropic Messages stream=true, aggregated into one Runtime artifact."""
|
|
client = await self._get_client()
|
|
full_response = ""
|
|
usage: Dict[str, int] = {}
|
|
response_id: Optional[str] = None
|
|
request_id: Optional[str] = None
|
|
try:
|
|
async with client.stream(
|
|
"POST",
|
|
self.llm_config.messages_endpoint,
|
|
json=self._anthropic_payload(conversation, stream=True),
|
|
) as response:
|
|
try:
|
|
response.raise_for_status()
|
|
except httpx.HTTPStatusError as exc:
|
|
self._raise_gateway_error(exc)
|
|
request_id = self._request_id_from_response(response)
|
|
async for line in response.aiter_lines():
|
|
if not line.startswith("data: "):
|
|
continue
|
|
data = line[6:]
|
|
if data == "[DONE]":
|
|
break
|
|
try:
|
|
event = json.loads(data)
|
|
except json.JSONDecodeError:
|
|
continue
|
|
event_type = event.get("type")
|
|
if event_type == "message_start":
|
|
message = event.get("message") or {}
|
|
response_id = response_id or message.get("id")
|
|
usage = self._normalize_usage(message.get("usage"))
|
|
elif event_type == "content_block_delta":
|
|
delta = event.get("delta") or {}
|
|
text = delta.get("text", "")
|
|
if text:
|
|
full_response += text
|
|
elif event_type == "message_delta":
|
|
delta_usage = (event.get("usage") or {})
|
|
if delta_usage:
|
|
usage = self._normalize_usage({**usage, **delta_usage})
|
|
|
|
conversation.add_message("assistant", full_response)
|
|
logger.info("收到 Claude Messages 流式回复", length=len(full_response), request_id=request_id)
|
|
return ModelResult(
|
|
content=full_response,
|
|
usage=usage,
|
|
request_id=request_id or response_id,
|
|
response_id=response_id,
|
|
model=self.llm_config.model,
|
|
api_format="anthropic_messages",
|
|
endpoint=self.llm_config.messages_endpoint,
|
|
)
|
|
except Exception as e:
|
|
logger.error("Claude Messages 流式请求失败", error=str(e))
|
|
raise
|
|
|
|
async def _stream_anthropic_messages_text(self, conversation: Conversation) -> AsyncGenerator[str, None]:
|
|
"""Yield text deltas from Anthropic Messages stream for A2A stream clients."""
|
|
client = await self._get_client()
|
|
full_response = ""
|
|
try:
|
|
async with client.stream(
|
|
"POST",
|
|
self.llm_config.messages_endpoint,
|
|
json=self._anthropic_payload(conversation, stream=True),
|
|
) as response:
|
|
try:
|
|
response.raise_for_status()
|
|
except httpx.HTTPStatusError as exc:
|
|
self._raise_gateway_error(exc)
|
|
async for line in response.aiter_lines():
|
|
if not line.startswith("data: "):
|
|
continue
|
|
data = line[6:]
|
|
if data == "[DONE]":
|
|
break
|
|
try:
|
|
event = json.loads(data)
|
|
except json.JSONDecodeError:
|
|
continue
|
|
if event.get("type") != "content_block_delta":
|
|
continue
|
|
delta = event.get("delta") or {}
|
|
text = delta.get("text", "")
|
|
if text:
|
|
full_response += text
|
|
yield text
|
|
conversation.add_message("assistant", full_response)
|
|
except Exception as e:
|
|
logger.error("Claude Messages 文本流失败", error=str(e))
|
|
raise
|
|
|
|
async def _stream_chat(self, conversation: Conversation) -> AsyncGenerator[str, None]:
|
|
"""流式对话"""
|
|
client = await self._get_client()
|
|
|
|
request_body = {
|
|
"model": self.llm_config.model,
|
|
"messages": conversation.to_openai_format(),
|
|
"temperature": self.llm_config.temperature,
|
|
"max_tokens": self.llm_config.max_tokens,
|
|
"stream": True,
|
|
"stream_options": {"include_usage": True},
|
|
}
|
|
|
|
full_response = ""
|
|
|
|
try:
|
|
async with client.stream(
|
|
"POST",
|
|
self.llm_config.chat_endpoint,
|
|
json=request_body
|
|
) as response:
|
|
response.raise_for_status()
|
|
|
|
async for line in response.aiter_lines():
|
|
if line.startswith("data: "):
|
|
data = line[6:]
|
|
if data == "[DONE]":
|
|
break
|
|
|
|
try:
|
|
chunk = json.loads(data)
|
|
choices = chunk.get("choices") or []
|
|
if not choices:
|
|
continue
|
|
delta = choices[0].get("delta", {})
|
|
content = delta.get("content", "")
|
|
if content:
|
|
full_response += content
|
|
yield content
|
|
except json.JSONDecodeError:
|
|
continue
|
|
|
|
# 保存完整回复到对话
|
|
conversation.add_message("assistant", full_response)
|
|
|
|
except Exception as e:
|
|
logger.error("流式请求失败", error=str(e))
|
|
raise
|
|
|
|
async def invoke_tool(self, tool_name: str, **kwargs) -> Any:
|
|
"""调用注册的工具"""
|
|
if tool_name not in self.tools:
|
|
raise ValueError(f"未找到工具: {tool_name}")
|
|
|
|
tool = self.tools[tool_name]
|
|
func = tool["function"]
|
|
|
|
logger.info(f"调用工具: {tool_name}", kwargs=kwargs)
|
|
|
|
if asyncio.iscoroutinefunction(func):
|
|
return await func(**kwargs)
|
|
else:
|
|
return func(**kwargs)
|
|
|
|
|
|
# 示例工具函数
|
|
def tool_get_current_time() -> str:
|
|
"""获取当前时间"""
|
|
return datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
|
|
|
|
|
def tool_calculate(expression: str) -> str:
|
|
"""计算数学表达式"""
|
|
try:
|
|
# 安全的数学计算
|
|
allowed_chars = set("0123456789+-*/.() ")
|
|
if not all(c in allowed_chars for c in expression):
|
|
return "错误: 不支持的字符"
|
|
result = eval(expression)
|
|
return str(result)
|
|
except Exception as e:
|
|
return f"计算错误: {str(e)}"
|