Coverage for agentos/llm/openai_provider.py: 25%
103 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 23:17 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 23:17 +0800
1"""
2OpenAI Provider 实现 — 基于官方 openai SDK 的对话补全。
3v1.3.36: +Function Calling / Tool Use 支持。
4"""
6from __future__ import annotations
8from collections.abc import Iterator
9from typing import Any
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
17from agentos.llm.base import (
18 CompletionChoice,
19 CompletionResult,
20 CompletionUsage,
21 LLMProvider,
22 Message,
23 MessageRole,
24 StreamChunk,
25 Tool,
26 ToolCall,
27)
29__all__ = ["OpenAIProvider"]
32_ROLE_MAP: dict[MessageRole, str] = {
33 MessageRole.SYSTEM: "system",
34 MessageRole.USER: "user",
35 MessageRole.ASSISTANT: "assistant",
36 MessageRole.TOOL: "tool",
37}
39_REVERSE_ROLE_MAP: dict[str, MessageRole] = {v: k for k, v in _ROLE_MAP.items()}
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}
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
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]
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
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 )
127class OpenAIProvider(LLMProvider):
128 """OpenAI SDK 提供商。支持 openai、azure、及所有 OpenAI 兼容的三方端点。"""
130 _sync_client: OpenAI | None = None
131 _async_client: AsyncOpenAI | None = None
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
145 @property
146 def provider_name(self) -> str:
147 return "openai"
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
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
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)
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)
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 )