Coverage for agentos/llm/base.py: 75%

108 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-09 07:12 +0800

1""" 

2LLM Provider 抽象层。 

3为 Nexus AgentOS 提供统一的 LLM 调用接口,实现 Provider 无关性。 

4v1.3.36: +Function Calling / Tool Use 抽象。 

5""" 

6 

7from __future__ import annotations 

8 

9from abc import ABC, abstractmethod 

10from collections.abc import Iterator 

11from dataclasses import dataclass, field 

12from enum import StrEnum 

13from typing import Any 

14 

15__all__ = [ 

16 "MessageRole", 

17 "Message", 

18 "CompletionUsage", 

19 "CompletionChoice", 

20 "CompletionResult", 

21 "StreamChunk", 

22 "TokenUsage", 

23 "Tool", 

24 "ToolFunction", 

25 "ToolParameter", 

26 "ToolCall", 

27 "LLMProvider", 

28] 

29 

30 

31class MessageRole(StrEnum): 

32 SYSTEM = "system" 

33 USER = "user" 

34 ASSISTANT = "assistant" 

35 TOOL = "tool" 

36 

37 

38@dataclass 

39class TokenUsage: 

40 prompt_tokens: int = 0 

41 completion_tokens: int = 0 

42 total_tokens: int = 0 

43 

44 

45@dataclass 

46class CompletionUsage(TokenUsage): 

47 cost_usd: float = 0.0 

48 

49 

50@dataclass 

51class Message: 

52 role: MessageRole 

53 content: str 

54 name: str | None = None 

55 tool_call_id: str | None = None 

56 tool_calls: list[ToolCall] | None = None 

57 

58 def as_dict(self) -> dict[str, Any]: 

59 d: dict[str, Any] = {"role": self.role.value, "content": self.content} 

60 if self.name: 

61 d["name"] = self.name 

62 if self.tool_call_id: 

63 d["tool_call_id"] = self.tool_call_id 

64 return d 

65 

66 

67# --- Function Calling / Tool Use --- 

68 

69 

70@dataclass 

71class ToolParameter: 

72 """JSON Schema 属性定义。""" 

73 

74 type: str = "string" 

75 description: str = "" 

76 enum: list[str] | None = None 

77 required: bool = False 

78 

79 def as_schema(self) -> dict[str, Any]: 

80 s: dict[str, Any] = {"type": self.type} 

81 if self.description: 

82 s["description"] = self.description 

83 if self.enum: 

84 s["enum"] = self.enum 

85 return s 

86 

87 

88@dataclass 

89class ToolFunction: 

90 """函数定义。""" 

91 

92 name: str 

93 description: str = "" 

94 parameters: dict[str, ToolParameter] = field(default_factory=dict) 

95 required: list[str] = field(default_factory=list) 

96 

97 def as_schema(self) -> dict[str, Any]: 

98 props = {k: v.as_schema() for k, v in self.parameters.items()} 

99 return { 

100 "type": "function", 

101 "function": { 

102 "name": self.name, 

103 "description": self.description, 

104 "parameters": { 

105 "type": "object", 

106 "properties": props, 

107 "required": self.required 

108 or [k for k, v in self.parameters.items() if v.required], 

109 }, 

110 }, 

111 } 

112 

113 

114@dataclass 

115class Tool: 

116 """顶层 Tool 包装。""" 

117 

118 function: ToolFunction 

119 

120 def as_schema(self) -> dict[str, Any]: 

121 return self.function.as_schema() 

122 

123 @classmethod 

124 def from_function( 

125 cls, 

126 name: str, 

127 description: str = "", 

128 parameters: dict[str, ToolParameter] | None = None, 

129 required: list[str] | None = None, 

130 ) -> Tool: 

131 return cls( 

132 function=ToolFunction( 

133 name=name, 

134 description=description, 

135 parameters=parameters or {}, 

136 required=required or [], 

137 ) 

138 ) 

139 

140 

141@dataclass 

142class ToolCall: 

143 """模型请求的工具调用。""" 

144 

145 id: str 

146 name: str 

147 arguments: str # JSON string 

148 

149 @property 

150 def parsed_arguments(self) -> dict[str, Any]: 

151 import json 

152 

153 return json.loads(self.arguments) 

154 

155 

156@dataclass 

157class CompletionChoice: 

158 index: int 

159 message: Message 

160 finish_reason: str = "stop" 

161 

162 

163@dataclass 

164class CompletionResult: 

165 id: str = "" 

166 model: str = "" 

167 choices: list[CompletionChoice] = field(default_factory=list) 

168 usage: CompletionUsage = field(default_factory=CompletionUsage) 

169 created: int = 0 

170 

171 

172@dataclass 

173class StreamChunk: 

174 content: str = "" 

175 finish_reason: str | None = None 

176 index: int = 0 

177 tool_calls: list[ToolCall] | None = None 

178 

179 

180class LLMProvider(ABC): 

181 """统一 LLM Provider 抽象。实现 OpenAI / Anthropic / 本地模型 的标准化调用。""" 

182 

183 def __init__(self, model: str = "", api_key: str = "", base_url: str = ""): 

184 self.model = model 

185 self.api_key = api_key 

186 self.base_url = base_url 

187 

188 @abstractmethod 

189 def chat( 

190 self, 

191 messages: list[Message], 

192 *, 

193 temperature: float = 0.7, 

194 max_tokens: int = 4096, 

195 top_p: float = 1.0, 

196 stop: list[str] | None = None, 

197 tools: list[Tool] | None = None, 

198 tool_choice: str = "auto", 

199 **kwargs: Any, 

200 ) -> CompletionResult: 

201 """同步聊天补全。""" 

202 ... 

203 

204 @abstractmethod 

205 async def achat( 

206 self, 

207 messages: list[Message], 

208 *, 

209 temperature: float = 0.7, 

210 max_tokens: int = 4096, 

211 top_p: float = 1.0, 

212 stop: list[str] | None = None, 

213 tools: list[Tool] | None = None, 

214 tool_choice: str = "auto", 

215 **kwargs: Any, 

216 ) -> CompletionResult: 

217 """异步聊天补全。""" 

218 ... 

219 

220 def stream( 

221 self, 

222 messages: list[Message], 

223 *, 

224 temperature: float = 0.7, 

225 max_tokens: int = 4096, 

226 tools: list[Tool] | None = None, 

227 **kwargs: Any, 

228 ) -> Iterator[StreamChunk]: 

229 """流式聊天补全。默认调用非流式包装。""" 

230 result = self.chat( 

231 messages, temperature=temperature, max_tokens=max_tokens, tools=tools, **kwargs 

232 ) 

233 for c in result.choices: 

234 yield StreamChunk( 

235 content=c.message.content, finish_reason=c.finish_reason, index=c.index 

236 ) 

237 

238 async def astream( 

239 self, 

240 messages: list[Message], 

241 *, 

242 temperature: float = 0.7, 

243 max_tokens: int = 4096, 

244 tools: list[Tool] | None = None, 

245 **kwargs: Any, 

246 ): 

247 """异步流式补全。默认调用 achat 包装。""" 

248 result = await self.achat( 

249 messages, temperature=temperature, max_tokens=max_tokens, tools=tools, **kwargs 

250 ) 

251 for c in result.choices: 

252 yield StreamChunk( 

253 content=c.message.content, finish_reason=c.finish_reason, index=c.index 

254 ) 

255 

256 @property 

257 @abstractmethod 

258 def provider_name(self) -> str: ...