Coverage for agentos/llm/openai_provider.py: 25%

103 statements  

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

1""" 

2OpenAI Provider 实现 — 基于官方 openai SDK 的对话补全。 

3v1.3.36: +Function Calling / Tool Use 支持。 

4""" 

5 

6from __future__ import annotations 

7 

8from collections.abc import Iterator 

9from typing import Any 

10 

11try: 

12 from openai import AsyncOpenAI, OpenAI 

13 from openai.types.chat import ChatCompletionMessageParam 

14except ImportError as e: 

15 raise ImportError("openai SDK not installed. Run: pip install 'nexus-agentos[openai]'") from e 

16 

17from agentos.llm.base import ( 

18 CompletionChoice, 

19 CompletionResult, 

20 CompletionUsage, 

21 LLMProvider, 

22 Message, 

23 MessageRole, 

24 StreamChunk, 

25 Tool, 

26 ToolCall, 

27) 

28 

29__all__ = ["OpenAIProvider"] 

30 

31 

32_ROLE_MAP: dict[MessageRole, str] = { 

33 MessageRole.SYSTEM: "system", 

34 MessageRole.USER: "user", 

35 MessageRole.ASSISTANT: "assistant", 

36 MessageRole.TOOL: "tool", 

37} 

38 

39_REVERSE_ROLE_MAP: dict[str, MessageRole] = {v: k for k, v in _ROLE_MAP.items()} 

40 

41# USD per 1K tokens (as of 2025-06) 

42_PRICING: dict[str, tuple[float, float]] = { 

43 "gpt-4o": (0.0025, 0.0100), 

44 "gpt-4o-mini": (0.00015, 0.0006), 

45 "gpt-4.1": (0.0020, 0.0080), 

46 "gpt-4.1-mini": (0.0004, 0.0016), 

47 "gpt-4.1-nano": (0.0001, 0.0004), 

48 "o3": (0.0100, 0.0400), 

49 "o3-mini": (0.0011, 0.0044), 

50 "o4-mini": (0.0011, 0.0044), 

51} 

52 

53 

54def _messages_to_openai(messages: list[Message]) -> list[ChatCompletionMessageParam]: 

55 """将 Message 列表转换为 OpenAI SDK 格式。""" 

56 result: list[ChatCompletionMessageParam] = [] 

57 for m in messages: 

58 entry: dict[str, Any] = {"role": _ROLE_MAP[m.role], "content": m.content} 

59 if m.tool_call_id: 

60 entry["tool_call_id"] = m.tool_call_id 

61 if m.tool_calls: 

62 entry["tool_calls"] = [ 

63 { 

64 "id": tc.id, 

65 "type": "function", 

66 "function": {"name": tc.name, "arguments": tc.arguments}, 

67 } 

68 for tc in m.tool_calls 

69 ] 

70 result.append(entry) 

71 return result 

72 

73 

74def _tools_to_openai(tools: list[Tool] | None) -> list[dict[str, Any]] | None: 

75 if not tools: 

76 return None 

77 return [t.as_schema() for t in tools] 

78 

79 

80def _extract_tool_calls(message_obj) -> list[ToolCall]: 

81 """从 OpenAI message 对象中提取 ToolCall 列表。""" 

82 raw = getattr(message_obj, "tool_calls", None) or [] 

83 result: list[ToolCall] = [] 

84 for tc in raw: 

85 fn = getattr(tc, "function", None) 

86 result.append( 

87 ToolCall( 

88 id=tc.id, 

89 name=fn.name if fn else "", 

90 arguments=fn.arguments if fn else "{}", 

91 ) 

92 ) 

93 return result 

94 

95 

96def _build_result(raw, model: str | None = None) -> CompletionResult: 

97 """从 OpenAI SDK 响应构建 CompletionResult。""" 

98 m = raw.choices[0].message 

99 role = _REVERSE_ROLE_MAP.get(m.role, MessageRole.ASSISTANT) 

100 tool_calls = _extract_tool_calls(m) 

101 choice = CompletionChoice( 

102 index=raw.choices[0].index, 

103 message=Message( 

104 role=role, 

105 content=m.content or "", 

106 tool_calls=tool_calls if tool_calls else None, 

107 ), 

108 finish_reason=raw.choices[0].finish_reason or "stop", 

109 ) 

110 usage = raw.usage 

111 tokens = CompletionUsage( 

112 prompt_tokens=usage.prompt_tokens if usage else 0, 

113 completion_tokens=usage.completion_tokens if usage else 0, 

114 total_tokens=usage.total_tokens if usage else 0, 

115 ) 

116 resolved_model = model or raw.model or "" 

117 if resolved_model in _PRICING: 

118 in_price, out_price = _PRICING[resolved_model] 

119 tokens.cost_usd = round( 

120 tokens.prompt_tokens / 1000 * in_price + tokens.completion_tokens / 1000 * out_price, 6 

121 ) 

122 return CompletionResult( 

123 id=raw.id, model=resolved_model, choices=[choice], usage=tokens, created=raw.created 

124 ) 

125 

126 

127class OpenAIProvider(LLMProvider): 

128 """OpenAI SDK 提供商。支持 openai、azure、及所有 OpenAI 兼容的三方端点。""" 

129 

130 _sync_client: OpenAI | None = None 

131 _async_client: AsyncOpenAI | None = None 

132 

133 def __init__( 

134 self, 

135 model: str = "gpt-4o-mini", 

136 api_key: str = "", 

137 base_url: str = "", 

138 organization: str = "", 

139 timeout: float = 60.0, 

140 ): 

141 super().__init__(model=model, api_key=api_key, base_url=base_url) 

142 self._organization = organization 

143 self._timeout = timeout 

144 

145 @property 

146 def provider_name(self) -> str: 

147 return "openai" 

148 

149 def _get_client(self) -> OpenAI: 

150 if self._sync_client is None: 

151 kwargs: dict[str, Any] = {"timeout": self._timeout, "max_retries": 2} 

152 if self.api_key: 

153 kwargs["api_key"] = self.api_key 

154 if self.base_url: 

155 kwargs["base_url"] = self.base_url 

156 if self._organization: 

157 kwargs["organization"] = self._organization 

158 self._sync_client = OpenAI(**kwargs) 

159 return self._sync_client 

160 

161 def _get_async_client(self) -> AsyncOpenAI: 

162 if self._async_client is None: 

163 kwargs: dict[str, Any] = {"timeout": self._timeout, "max_retries": 2} 

164 if self.api_key: 

165 kwargs["api_key"] = self.api_key 

166 if self.base_url: 

167 kwargs["base_url"] = self.base_url 

168 if self._organization: 

169 kwargs["organization"] = self._organization 

170 self._async_client = AsyncOpenAI(**kwargs) 

171 return self._async_client 

172 

173 def chat( 

174 self, 

175 messages: list[Message], 

176 *, 

177 temperature: float = 0.7, 

178 max_tokens: int = 4096, 

179 top_p: float = 1.0, 

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

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

182 tool_choice: str = "auto", 

183 **kwargs: Any, 

184 ) -> CompletionResult: 

185 client = self._get_client() 

186 params: dict[str, Any] = { 

187 "model": self.model, 

188 "messages": _messages_to_openai(messages), 

189 "temperature": temperature, 

190 "max_tokens": max_tokens, 

191 "top_p": top_p, 

192 "stop": stop, 

193 **kwargs, 

194 } 

195 if tools: 

196 params["tools"] = _tools_to_openai(tools) 

197 params["tool_choice"] = tool_choice 

198 resp = client.chat.completions.create(**params) 

199 return _build_result(resp, model=self.model) 

200 

201 async def achat( 

202 self, 

203 messages: list[Message], 

204 *, 

205 temperature: float = 0.7, 

206 max_tokens: int = 4096, 

207 top_p: float = 1.0, 

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

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

210 tool_choice: str = "auto", 

211 **kwargs: Any, 

212 ) -> CompletionResult: 

213 client = self._get_async_client() 

214 params: dict[str, Any] = { 

215 "model": self.model, 

216 "messages": _messages_to_openai(messages), 

217 "temperature": temperature, 

218 "max_tokens": max_tokens, 

219 "top_p": top_p, 

220 "stop": stop, 

221 **kwargs, 

222 } 

223 if tools: 

224 params["tools"] = _tools_to_openai(tools) 

225 params["tool_choice"] = tool_choice 

226 resp = await client.chat.completions.create(**params) 

227 return _build_result(resp, model=self.model) 

228 

229 def stream( 

230 self, 

231 messages: list[Message], 

232 *, 

233 temperature: float = 0.7, 

234 max_tokens: int = 4096, 

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

236 **kwargs: Any, 

237 ) -> Iterator[StreamChunk]: 

238 client = self._get_client() 

239 params: dict[str, Any] = { 

240 "model": self.model, 

241 "messages": _messages_to_openai(messages), 

242 "temperature": temperature, 

243 "max_tokens": max_tokens, 

244 "stream": True, 

245 **kwargs, 

246 } 

247 if tools: 

248 params["tools"] = _tools_to_openai(tools) 

249 stream_resp = client.chat.completions.create(**params) 

250 for chunk in stream_resp: 

251 if chunk.choices and chunk.choices[0].delta.content: 

252 yield StreamChunk( 

253 content=chunk.choices[0].delta.content, 

254 finish_reason=( 

255 chunk.choices[0].finish_reason if chunk.choices[0].finish_reason else None 

256 ), 

257 index=chunk.choices[0].index, 

258 )