""" 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)}"