Coverage for agentos/llm/openai_provider.py: 25%
102 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
1"""
2OpenAI Provider 实现 — 基于官方 openai SDK 的对话补全。
3v1.3.36: +Function Calling / Tool Use 支持。
4"""
6from __future__ import annotations
8from typing import Any, Iterator
10try:
11 from openai import AsyncOpenAI, OpenAI
12 from openai.types.chat import ChatCompletionMessageParam
13except ImportError as e:
14 raise ImportError(
15 "openai SDK not installed. Run: pip install 'nexus-agentos[openai]'"
16 ) from e
18from agentos.llm.base import (
19 CompletionChoice,
20 CompletionResult,
21 CompletionUsage,
22 LLMProvider,
23 Message,
24 MessageRole,
25 StreamChunk,
26 Tool,
27 ToolCall,
28)
30__all__ = ["OpenAIProvider"]
33_ROLE_MAP: dict[MessageRole, str] = {
34 MessageRole.SYSTEM: "system",
35 MessageRole.USER: "user",
36 MessageRole.ASSISTANT: "assistant",
37 MessageRole.TOOL: "tool",
38}
40_REVERSE_ROLE_MAP: dict[str, MessageRole] = {v: k for k, v in _ROLE_MAP.items()}
42# USD per 1K tokens (as of 2025-06)
43_PRICING: dict[str, tuple[float, float]] = {
44 "gpt-4o": (0.0025, 0.0100),
45 "gpt-4o-mini": (0.00015, 0.0006),
46 "gpt-4.1": (0.0020, 0.0080),
47 "gpt-4.1-mini": (0.0004, 0.0016),
48 "gpt-4.1-nano": (0.0001, 0.0004),
49 "o3": (0.0100, 0.0400),
50 "o3-mini": (0.0011, 0.0044),
51 "o4-mini": (0.0011, 0.0044),
52}
55def _messages_to_openai(messages: list[Message]) -> list[ChatCompletionMessageParam]:
56 """将 Message 列表转换为 OpenAI SDK 格式。"""
57 result: list[ChatCompletionMessageParam] = []
58 for m in messages:
59 entry: dict[str, Any] = {"role": _ROLE_MAP[m.role], "content": m.content}
60 if m.tool_call_id:
61 entry["tool_call_id"] = m.tool_call_id
62 if m.tool_calls:
63 entry["tool_calls"] = [
64 {
65 "id": tc.id,
66 "type": "function",
67 "function": {"name": tc.name, "arguments": tc.arguments},
68 }
69 for tc in m.tool_calls
70 ]
71 result.append(entry)
72 return result
75def _tools_to_openai(tools: list[Tool] | None) -> list[dict[str, Any]] | None:
76 if not tools:
77 return None
78 return [t.as_schema() for t in tools]
81def _extract_tool_calls(message_obj) -> list[ToolCall]:
82 """从 OpenAI message 对象中提取 ToolCall 列表。"""
83 raw = getattr(message_obj, "tool_calls", None) or []
84 result: list[ToolCall] = []
85 for tc in raw:
86 fn = getattr(tc, "function", None)
87 result.append(ToolCall(
88 id=tc.id,
89 name=fn.name if fn else "",
90 arguments=fn.arguments if fn else "{}",
91 ))
92 return result
95def _build_result(raw, model: str | None = None) -> CompletionResult:
96 """从 OpenAI SDK 响应构建 CompletionResult。"""
97 m = raw.choices[0].message
98 role = _REVERSE_ROLE_MAP.get(m.role, MessageRole.ASSISTANT)
99 tool_calls = _extract_tool_calls(m)
100 choice = CompletionChoice(
101 index=raw.choices[0].index,
102 message=Message(
103 role=role, content=m.content or "",
104 tool_calls=tool_calls if tool_calls else None,
105 ),
106 finish_reason=raw.choices[0].finish_reason or "stop",
107 )
108 usage = raw.usage
109 tokens = CompletionUsage(
110 prompt_tokens=usage.prompt_tokens if usage else 0,
111 completion_tokens=usage.completion_tokens if usage else 0,
112 total_tokens=usage.total_tokens if usage else 0,
113 )
114 resolved_model = model or raw.model or ""
115 if resolved_model in _PRICING:
116 in_price, out_price = _PRICING[resolved_model]
117 tokens.cost_usd = round(
118 tokens.prompt_tokens / 1000 * in_price + tokens.completion_tokens / 1000 * out_price, 6
119 )
120 return CompletionResult(
121 id=raw.id, model=resolved_model, choices=[choice], usage=tokens, created=raw.created
122 )
125class OpenAIProvider(LLMProvider):
126 """OpenAI SDK 提供商。支持 openai、azure、及所有 OpenAI 兼容的三方端点。"""
128 _sync_client: OpenAI | None = None
129 _async_client: AsyncOpenAI | None = None
131 def __init__(
132 self,
133 model: str = "gpt-4o-mini",
134 api_key: str = "",
135 base_url: str = "",
136 organization: str = "",
137 timeout: float = 60.0,
138 ):
139 super().__init__(model=model, api_key=api_key, base_url=base_url)
140 self._organization = organization
141 self._timeout = timeout
143 @property
144 def provider_name(self) -> str:
145 return "openai"
147 def _get_client(self) -> OpenAI:
148 if self._sync_client is None:
149 kwargs: dict[str, Any] = {"timeout": self._timeout, "max_retries": 2}
150 if self.api_key:
151 kwargs["api_key"] = self.api_key
152 if self.base_url:
153 kwargs["base_url"] = self.base_url
154 if self._organization:
155 kwargs["organization"] = self._organization
156 self._sync_client = OpenAI(**kwargs)
157 return self._sync_client
159 def _get_async_client(self) -> AsyncOpenAI:
160 if self._async_client is None:
161 kwargs: dict[str, Any] = {"timeout": self._timeout, "max_retries": 2}
162 if self.api_key:
163 kwargs["api_key"] = self.api_key
164 if self.base_url:
165 kwargs["base_url"] = self.base_url
166 if self._organization:
167 kwargs["organization"] = self._organization
168 self._async_client = AsyncOpenAI(**kwargs)
169 return self._async_client
171 def chat(
172 self,
173 messages: list[Message],
174 *,
175 temperature: float = 0.7,
176 max_tokens: int = 4096,
177 top_p: float = 1.0,
178 stop: list[str] | None = None,
179 tools: list[Tool] | None = None,
180 tool_choice: str = "auto",
181 **kwargs: Any,
182 ) -> CompletionResult:
183 client = self._get_client()
184 params: dict[str, Any] = {
185 "model": self.model,
186 "messages": _messages_to_openai(messages),
187 "temperature": temperature,
188 "max_tokens": max_tokens,
189 "top_p": top_p,
190 "stop": stop,
191 **kwargs,
192 }
193 if tools:
194 params["tools"] = _tools_to_openai(tools)
195 params["tool_choice"] = tool_choice
196 resp = client.chat.completions.create(**params)
197 return _build_result(resp, model=self.model)
199 async def achat(
200 self,
201 messages: list[Message],
202 *,
203 temperature: float = 0.7,
204 max_tokens: int = 4096,
205 top_p: float = 1.0,
206 stop: list[str] | None = None,
207 tools: list[Tool] | None = None,
208 tool_choice: str = "auto",
209 **kwargs: Any,
210 ) -> CompletionResult:
211 client = self._get_async_client()
212 params: dict[str, Any] = {
213 "model": self.model,
214 "messages": _messages_to_openai(messages),
215 "temperature": temperature,
216 "max_tokens": max_tokens,
217 "top_p": top_p,
218 "stop": stop,
219 **kwargs,
220 }
221 if tools:
222 params["tools"] = _tools_to_openai(tools)
223 params["tool_choice"] = tool_choice
224 resp = await client.chat.completions.create(**params)
225 return _build_result(resp, model=self.model)
227 def stream(
228 self,
229 messages: list[Message],
230 *,
231 temperature: float = 0.7,
232 max_tokens: int = 4096,
233 tools: list[Tool] | None = None,
234 **kwargs: Any,
235 ) -> Iterator[StreamChunk]:
236 client = self._get_client()
237 params: dict[str, Any] = {
238 "model": self.model,
239 "messages": _messages_to_openai(messages),
240 "temperature": temperature,
241 "max_tokens": max_tokens,
242 "stream": True,
243 **kwargs,
244 }
245 if tools:
246 params["tools"] = _tools_to_openai(tools)
247 stream_resp = client.chat.completions.create(**params)
248 for chunk in stream_resp:
249 if chunk.choices and chunk.choices[0].delta.content:
250 yield StreamChunk(
251 content=chunk.choices[0].delta.content,
252 finish_reason=(
253 chunk.choices[0].finish_reason if chunk.choices[0].finish_reason else None
254 ),
255 index=chunk.choices[0].index,
256 )