-
Notifications
You must be signed in to change notification settings - Fork 64
Expand file tree
/
Copy pathopenai.py
More file actions
308 lines (270 loc) · 12.2 KB
/
Copy pathopenai.py
File metadata and controls
308 lines (270 loc) · 12.2 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
"""OpenAI-compatible API backend (query + generate). Supports Qwen API and any OpenAI-compatible endpoint."""
import json
import logging
import re
import time
from typing import Any
from openai import OpenAI
from config import Config
from .gemini import FunctionSpec, compile_prompt_to_md
from .model_profiles import get_profile, supports_json_schema, thinking_json_incompatible, supports_tool_choice_required, get_thinking_extra_body
logger = logging.getLogger("MLEvolve")
def _strip_markdown_fences(args: str) -> str:
"""Remove markdown code fences that LLMs sometimes append inside JSON string values."""
cleaned = re.sub(r'\\n```[a-z]*\s*("?\s*\}?\s*)$', r'\1', args.rstrip())
cleaned = cleaned.rstrip()
if not cleaned.endswith('}'):
if not cleaned.endswith('"'):
cleaned += '"'
cleaned += '}'
return cleaned
def _parse_json_args(args: str) -> dict:
"""Parse function call arguments, tolerating Python literals and markdown fences."""
# 1. Fast path: valid JSON as-is
try:
return json.loads(args)
except json.JSONDecodeError:
pass
# 2. Try stripping markdown fences
try:
cleaned = _strip_markdown_fences(args)
if cleaned != args:
result = json.loads(cleaned)
logger.warning("Fixed malformed function args by stripping markdown code fences")
return result
except json.JSONDecodeError:
pass
# 3. Normalize Python literals (None/True/False) outside quoted strings
parts = re.split(r'("(?:[^"\\]|\\.)*")', args)
normalized = []
for part in parts:
if part.startswith('"'):
normalized.append(part)
else:
part = re.sub(r'\bNone\b', 'null', part)
part = re.sub(r'\bTrue\b', 'true', part)
part = re.sub(r'\bFalse\b', 'false', part)
normalized.append(part)
normalized_str = ''.join(normalized)
try:
return json.loads(normalized_str)
except json.JSONDecodeError:
pass
# 4. Normalized + strip markdown fences
cleaned = _strip_markdown_fences(normalized_str)
return json.loads(cleaned)
# Return type aligned with gemini.query
OutputType = str | dict
def _stage_config_for_model(cfg: Config, model: str):
"""Return code or feedback config depending on which model is being used."""
if cfg.agent.code.model == model:
return cfg.agent.code
return cfg.agent.feedback
def _build_messages(system_message: str | None, user_message: str | None, model: str = "") -> list[dict[str, str]]:
# Anthropic API (Claude) requires the messages array to contain at least
# one user-role message; system is a separate top-level field. When only
# system_message is provided, the OpenAI-compat proxy converts
# [{role: system}] -> system="...", messages=[] which Anthropic rejects
# with "field messages is required". Promote system to user in that case.
is_claude = (model or "").lower().startswith("claude")
if is_claude and system_message and not user_message:
return [{"role": "user", "content": system_message}]
messages = []
if system_message:
messages.append({"role": "system", "content": system_message})
if user_message:
messages.append({"role": "user", "content": user_message})
return messages
def query(
system_message: str | None,
user_message: str | None,
func_spec: FunctionSpec | None = None,
cfg: Config | None = None,
**model_kwargs,
) -> tuple[OutputType, float, int, int, dict]:
"""OpenAI-compatible query (chat completions, optional function calling). Same return shape as gemini.query."""
if cfg is None:
raise ValueError("cfg is required for OpenAI backend")
filtered = {k: v for k, v in model_kwargs.items() if v is not None}
model = filtered.get("model", "")
stage = _stage_config_for_model(cfg, model)
client = OpenAI(
api_key=stage.api_key,
base_url=stage.base_url or None,
timeout=1200.0,
)
messages = _build_messages(system_message, user_message, model=model)
if not messages:
raise ValueError("Either system_message or user_message must be provided")
# Function calling requires non_thinking mode for Qwen (errors on
# tool_choice=required + thinking). Claude supports thinking + tool use
# as long as tool_choice is auto/none — handled below by
# _NO_TOOL_CHOICE_REQUIRED_PREFIXES, which keeps tool_choice=auto for Claude.
is_claude = model.lower().startswith("claude")
use_thinking = func_spec is None or is_claude
profile = get_profile(model, use_thinking=use_thinking)
extra_body: dict[str, Any] = {}
if "top_k" in profile:
extra_body["top_k"] = profile["top_k"]
if "enable_thinking" in profile:
extra_body["enable_thinking"] = profile["enable_thinking"]
# Merge model-specific thinking params (synced from agentic-mle)
if use_thinking:
extra_body.update(get_thinking_extra_body(model))
params: dict[str, Any] = {
"model": model,
"messages": messages,
"temperature": profile.get("temperature", filtered.get("temperature", 1.0)),
"max_tokens": filtered.get("max_tokens", 16384),
}
if "top_p" in profile:
params["top_p"] = profile["top_p"]
if "presence_penalty" in profile:
params["presence_penalty"] = profile["presence_penalty"]
if extra_body:
params["extra_body"] = extra_body
if func_spec is not None:
tool_dict = func_spec.as_openai_tool_dict
if not supports_json_schema(model):
tool_dict.pop("strict", None)
params["tools"] = [tool_dict]
if supports_tool_choice_required(model):
params["tool_choice"] = func_spec.openai_tool_choice_dict
t0 = time.time()
logger.info(f"Querying OpenAI-compatible API with model: {model}")
try:
completion = client.chat.completions.create(**params)
except Exception as e:
logger.error(f"Error calling OpenAI-compatible API: {e}")
raise
req_time = time.time() - t0
choice = completion.choices[0]
message = choice.message
if getattr(choice, "finish_reason", None) == "length":
logger.warning(f"Response truncated by max_tokens ({params.get('max_tokens')}), consider increasing it")
if func_spec is None:
output = message.content or ""
logger.info(f"OpenAI response: {output}", extra={"verbose": True})
else:
if not message.tool_calls:
raise ValueError("Expected function call, got no tool_calls")
tc = message.tool_calls[0]
if tc.function.name != func_spec.name:
raise ValueError(f"Function name mismatch: expected {func_spec.name}, got {tc.function.name}")
try:
output = _parse_json_args(tc.function.arguments or "{}")
except json.JSONDecodeError as e:
logger.error(f"Invalid function arguments: {tc.function.arguments}")
raise e
logger.info(f"OpenAI function call response: {output}", extra={"verbose": True})
in_tok = getattr(completion.usage, "prompt_tokens", 0) or 0
out_tok = getattr(completion.usage, "completion_tokens", 0) or 0
info = {
"model": getattr(completion, "model", model),
"created": getattr(completion, "created", int(time.time())),
}
return output, req_time, in_tok, out_tok, info
def _prompt_to_messages(prompt: str | dict | list, model: str = "") -> list[dict[str, str]]:
"""Convert prompt to chat messages. Supports Qwen/OpenAI chat format: {system, user, assistant}.
For GPT models, assistant content is appended to the user message instead of
being sent as a separate assistant message, because GPT models may return
empty responses when they see a trailing assistant prefill.
"""
if isinstance(prompt, dict) and ("system" in prompt or "user" in prompt or "assistant" in prompt):
messages = []
if prompt.get("system"):
messages.append({"role": "system", "content": str(prompt["system"])})
is_gpt = (model or "").lower().startswith("gpt")
user_content = str(prompt["user"]) if prompt.get("user") else ""
assistant_content = str(prompt["assistant"]) if prompt.get("assistant") else ""
if is_gpt and assistant_content:
# GPT: merge assistant prefill into user message
combined = f"{user_content}\n\n{assistant_content}" if user_content else assistant_content
messages.append({"role": "user", "content": combined})
else:
if user_content:
messages.append({"role": "user", "content": user_content})
if assistant_content:
messages.append({"role": "assistant", "content": assistant_content})
if not messages:
raise ValueError("Chat dict must have at least one of: system, user, assistant")
return messages
content = prompt if isinstance(prompt, str) else compile_prompt_to_md(prompt)
return [{"role": "user", "content": content}]
def generate(
prompt: str | dict | list,
cfg: Config,
temperature: float | None = None,
max_tokens: int | None = None,
stop_tokens: list[str] | None = None,
json_schema: dict | None = None,
max_retries: int = 20,
retry_delay: float = 3,
) -> str:
"""Streaming text generation via OpenAI-compatible Chat API. Supports chat format {system, user, assistant} for Qwen."""
stage = cfg.agent.code
model = stage.model
messages = _prompt_to_messages(prompt, model=model)
client = OpenAI(
api_key=stage.api_key,
base_url=stage.base_url or None,
timeout=1200.0,
)
# Qwen: thinking + json_schema are mutually exclusive — drop schema, keep thinking.
if json_schema is not None and thinking_json_incompatible(model):
json_schema = None
# Claude: adaptive thinking + json_schema both supported — always keep thinking ON.
# Other models: thinking on only when no json_schema (legacy Qwen-aligned behavior).
is_claude = model.lower().startswith("claude")
use_thinking = json_schema is None or is_claude
profile = get_profile(model, use_thinking=use_thinking)
extra_body: dict[str, Any] = {}
if "top_k" in profile:
extra_body["top_k"] = profile["top_k"]
if "enable_thinking" in profile:
extra_body["enable_thinking"] = profile["enable_thinking"]
# Merge model-specific thinking params (synced from agentic-mle)
if use_thinking:
extra_body.update(get_thinking_extra_body(model))
params: dict[str, Any] = {
"model": model,
"messages": messages,
"temperature": profile.get("temperature", temperature if temperature is not None else 1.0),
"max_tokens": max_tokens if max_tokens is not None else 16384,
"stream": True,
}
if "top_p" in profile:
params["top_p"] = profile["top_p"]
if "presence_penalty" in profile:
params["presence_penalty"] = profile["presence_penalty"]
if extra_body:
params["extra_body"] = extra_body
if stop_tokens:
params["stop"] = stop_tokens
if json_schema is not None:
if supports_json_schema(model):
params["response_format"] = {
"type": "json_schema",
"json_schema": {"name": "structured_output", "strict": False, "schema": json_schema},
}
else:
params["response_format"] = {"type": "json_object"}
logger.info(f"generate messages: {len(messages)} turns", extra={"verbose": True})
for attempt in range(max_retries):
try:
stream = client.chat.completions.create(**params)
full_text = ""
for chunk in stream:
if chunk.choices and chunk.choices[0].delta.content:
full_text += chunk.choices[0].delta.content
if "</think>" in full_text:
full_text = full_text[full_text.find("</think>") + 8:]
logger.info(f"generate response: {full_text}", extra={"verbose": True})
return full_text
except Exception as e:
logger.warning(f"generate failed, retrying {attempt + 1}/{max_retries}: {e}")
if attempt >= max_retries - 1:
logger.error("generate retry limit reached")
raise
time.sleep(retry_delay)
return ""