341 lines
12 KiB
Python
341 lines
12 KiB
Python
"""
|
|
OpenAI API 客户端封装(同步/异步)
|
|
|
|
提供与 :mod:`llm.deepseek_client` 同形态的接口,便于在需要时切换 LLM 路由。
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import os
|
|
from dataclasses import dataclass
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
from openai import AsyncOpenAI, OpenAI # type: ignore[import-not-found]
|
|
from openai.types.chat import ( # type: ignore[import-not-found]
|
|
ChatCompletion,
|
|
ChatCompletionMessage,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
@dataclass
|
|
class OpenAIConfig:
|
|
"""OpenAI API 配置(支持 OpenAI 兼容网关)。"""
|
|
|
|
api_key: str
|
|
base_url: str = "https://api.openai.com/v1"
|
|
model_name: str = "gpt-4o-mini"
|
|
temperature: float = 0.3
|
|
max_tokens: int = 4096
|
|
top_p: float = 0.9
|
|
frequency_penalty: float = 0.0
|
|
presence_penalty: float = 0.0
|
|
timeout: float = 60.0
|
|
extra_headers: Optional[Dict[str, str]] = None
|
|
|
|
|
|
class OpenAIClient:
|
|
"""
|
|
OpenAI API 客户端(同步)。
|
|
|
|
说明:本项目其他模块仅要求具备 ``chat`` / ``chat_with_json`` 方法;
|
|
这里额外提供若干便捷方法,与 DeepSeekClient 对齐,便于复用同一套 prompt。
|
|
"""
|
|
|
|
def __init__(self, config: OpenAIConfig):
|
|
self.config = config
|
|
self.client = OpenAI(
|
|
api_key=config.api_key,
|
|
base_url=config.base_url,
|
|
timeout=config.timeout,
|
|
)
|
|
logger.info(
|
|
"[OK] OpenAIClient初始化: model=%s, base_url=%s",
|
|
config.model_name,
|
|
config.base_url,
|
|
)
|
|
|
|
def chat(self, messages: List[Dict[str, str]], **kwargs) -> ChatCompletionMessage:
|
|
max_tokens = kwargs.get("max_tokens", self.config.max_tokens)
|
|
max_completion_tokens = kwargs.get("max_completion_tokens", None)
|
|
|
|
params: Dict[str, Any] = {
|
|
"model": self.config.model_name,
|
|
"messages": messages,
|
|
"temperature": kwargs.get("temperature", self.config.temperature),
|
|
"top_p": kwargs.get("top_p", self.config.top_p),
|
|
"frequency_penalty": kwargs.get(
|
|
"frequency_penalty", self.config.frequency_penalty
|
|
),
|
|
"presence_penalty": kwargs.get(
|
|
"presence_penalty", self.config.presence_penalty
|
|
),
|
|
"timeout": kwargs.get("timeout", self.config.timeout),
|
|
}
|
|
|
|
# 兼容不同模型/网关:
|
|
# - 传统 Chat Completions 使用 max_tokens
|
|
# - 部分新模型(如 gpt-5.*)要求 max_completion_tokens
|
|
if max_completion_tokens is not None:
|
|
params["max_completion_tokens"] = max_completion_tokens
|
|
else:
|
|
params["max_tokens"] = max_tokens
|
|
|
|
if self.config.extra_headers:
|
|
params["extra_headers"] = self.config.extra_headers
|
|
|
|
try:
|
|
response: ChatCompletion = self.client.chat.completions.create(**params)
|
|
message = response.choices[0].message
|
|
usage = response.usage
|
|
if usage is not None:
|
|
logger.debug(
|
|
"OpenAI调用完成: prompt_tokens=%s, completion_tokens=%s, total_tokens=%s",
|
|
usage.prompt_tokens,
|
|
usage.completion_tokens,
|
|
usage.total_tokens,
|
|
)
|
|
return message
|
|
except Exception as e:
|
|
# 自动兼容:若网关提示 max_tokens 不支持,则改用 max_completion_tokens 重试一次
|
|
msg = str(e)
|
|
if (
|
|
"Unsupported parameter" in msg
|
|
and "max_tokens" in msg
|
|
and "max_completion_tokens" in msg
|
|
and "max_completion_tokens" not in params
|
|
):
|
|
params.pop("max_tokens", None)
|
|
params["max_completion_tokens"] = max_tokens
|
|
try:
|
|
response = self.client.chat.completions.create(**params)
|
|
return response.choices[0].message
|
|
except Exception:
|
|
pass
|
|
|
|
logger.error("OpenAI API调用失败: %s", e)
|
|
raise
|
|
|
|
def chat_with_json(
|
|
self, messages: List[Dict[str, str]], **kwargs
|
|
) -> Dict[str, Any]:
|
|
message = self.chat(messages, **kwargs)
|
|
content = (message.content or "").strip()
|
|
|
|
try:
|
|
if "```json" in content:
|
|
start = content.find("```json") + 7
|
|
end = content.find("```", start)
|
|
content = content[start:end].strip()
|
|
elif "```" in content:
|
|
start = content.find("```") + 3
|
|
end = content.find("```", start)
|
|
content = content[start:end].strip()
|
|
return json.loads(content)
|
|
except json.JSONDecodeError as e:
|
|
logger.warning("JSON解析失败,返回原始内容: %s", e)
|
|
return {"_json_decode_failed": True, "raw_content": content}
|
|
|
|
# ===== 便捷方法(与 DeepSeekClient 对齐)=====
|
|
def generate_sql(
|
|
self, prompt: str, schema: str, dialect: str = "tsql", **kwargs
|
|
) -> str:
|
|
from config.prompts import SQL_GENERATOR_SYSTEM, SQL_GENERATOR_USER
|
|
|
|
messages = [
|
|
{"role": "system", "content": SQL_GENERATOR_SYSTEM},
|
|
{
|
|
"role": "user",
|
|
"content": SQL_GENERATOR_USER.format(
|
|
schema=schema, question=prompt, dialect=dialect
|
|
),
|
|
},
|
|
]
|
|
|
|
kwargs.setdefault("temperature", 0.0)
|
|
kwargs.setdefault("top_p", 1.0)
|
|
response = self.chat(messages, **kwargs)
|
|
content = (response.content or "").strip()
|
|
|
|
if "```sql" in content:
|
|
start = content.find("```sql") + 6
|
|
end = content.find("```", start)
|
|
content = content[start:end].strip()
|
|
elif "```" in content:
|
|
start = content.find("```") + 3
|
|
end = content.find("```", start)
|
|
content = content[start:end].strip()
|
|
|
|
return content
|
|
|
|
def validate_sql(self, sql: str, schema: str, **kwargs) -> Dict[str, Any]:
|
|
from config.prompts import VALIDATOR_SYSTEM, VALIDATOR_USER
|
|
|
|
messages = [
|
|
{"role": "system", "content": VALIDATOR_SYSTEM},
|
|
{"role": "user", "content": VALIDATOR_USER.format(sql=sql, schema=schema)},
|
|
]
|
|
|
|
kwargs.setdefault("temperature", 0.0)
|
|
kwargs.setdefault("top_p", 1.0)
|
|
return self.chat_with_json(messages, **kwargs)
|
|
|
|
def select_tables(
|
|
self, question: str, table_list: str, **kwargs
|
|
) -> Dict[str, Any]:
|
|
from config.prompts import SCHEMA_LINKER_SYSTEM, SCHEMA_LINKER_USER
|
|
|
|
messages = [
|
|
{"role": "system", "content": SCHEMA_LINKER_SYSTEM},
|
|
{
|
|
"role": "user",
|
|
"content": SCHEMA_LINKER_USER.format(
|
|
question=question, table_list=table_list
|
|
),
|
|
},
|
|
]
|
|
|
|
kwargs.setdefault("temperature", 0.0)
|
|
kwargs.setdefault("top_p", 1.0)
|
|
return self.chat_with_json(messages, **kwargs)
|
|
|
|
def normalize_nl_question_for_text2sql(self, question: str) -> str:
|
|
from config.prompts import (
|
|
CANONICALIZE_NL_FOR_TEXT2SQL_SYSTEM,
|
|
CANONICALIZE_NL_FOR_TEXT2SQL_USER,
|
|
)
|
|
|
|
q = (question or "").strip()
|
|
if not q:
|
|
return ""
|
|
messages = [
|
|
{"role": "system", "content": CANONICALIZE_NL_FOR_TEXT2SQL_SYSTEM},
|
|
{"role": "user", "content": CANONICALIZE_NL_FOR_TEXT2SQL_USER.format(question=q)},
|
|
]
|
|
msg = self.chat(messages, temperature=0.0, top_p=1.0, max_completion_tokens=512)
|
|
text = (msg.content or "").strip()
|
|
line = text.splitlines()[0].strip() if text else ""
|
|
return line.strip("「」\"'“”")
|
|
|
|
def translate_nl_question_to_zh(self, question: str) -> str:
|
|
return self.normalize_nl_question_for_text2sql(question)
|
|
|
|
def empty_result_user_feedback(
|
|
self,
|
|
question: str,
|
|
sql: str,
|
|
schema: str,
|
|
*,
|
|
max_schema_chars: int = 8000,
|
|
**kwargs: Any,
|
|
) -> str:
|
|
from config.prompts import EMPTY_RESULT_FEEDBACK_SYSTEM, EMPTY_RESULT_FEEDBACK_USER
|
|
|
|
schema_snip = (schema or "")[:max_schema_chars]
|
|
messages = [
|
|
{"role": "system", "content": EMPTY_RESULT_FEEDBACK_SYSTEM},
|
|
{
|
|
"role": "user",
|
|
"content": EMPTY_RESULT_FEEDBACK_USER.format(
|
|
question=question or "(无)",
|
|
sql=sql,
|
|
schema=schema_snip,
|
|
),
|
|
},
|
|
]
|
|
msg = self.chat(messages, temperature=0.4, max_completion_tokens=512, **kwargs)
|
|
return ((msg.content or "").strip())
|
|
|
|
def sql_probe_success_delivery_message(
|
|
self,
|
|
question: str,
|
|
sql: str,
|
|
*,
|
|
max_sql_chars: int = 4000,
|
|
**kwargs: Any,
|
|
) -> str:
|
|
from config.prompts import (
|
|
SQL_PROBE_SUCCESS_DELIVERY_SYSTEM,
|
|
SQL_PROBE_SUCCESS_DELIVERY_USER,
|
|
)
|
|
|
|
q = (question or "").strip() or "(无)"
|
|
s = (sql or "").strip()
|
|
if len(s) > max_sql_chars:
|
|
s = s[: max_sql_chars - 20].rstrip() + "\n-- …(已截断)"
|
|
|
|
messages = [
|
|
{"role": "system", "content": SQL_PROBE_SUCCESS_DELIVERY_SYSTEM},
|
|
{
|
|
"role": "user",
|
|
"content": SQL_PROBE_SUCCESS_DELIVERY_USER.format(question=q, sql=s),
|
|
},
|
|
]
|
|
kwargs.setdefault("temperature", 0.2)
|
|
kwargs.setdefault("top_p", 1.0)
|
|
kwargs.setdefault("max_completion_tokens", 320)
|
|
msg = self.chat(messages, **kwargs)
|
|
return (msg.content or "").strip()
|
|
|
|
|
|
class AsyncOpenAIClient:
|
|
"""OpenAI API 客户端(异步)。"""
|
|
|
|
def __init__(self, config: OpenAIConfig):
|
|
self.config = config
|
|
self.client = AsyncOpenAI(
|
|
api_key=config.api_key,
|
|
base_url=config.base_url,
|
|
timeout=config.timeout,
|
|
)
|
|
|
|
async def chat(self, messages: List[Dict[str, str]], **kwargs) -> ChatCompletionMessage:
|
|
params: Dict[str, Any] = {
|
|
"model": self.config.model_name,
|
|
"messages": messages,
|
|
"temperature": kwargs.get("temperature", self.config.temperature),
|
|
"max_tokens": kwargs.get("max_tokens", self.config.max_tokens),
|
|
"top_p": kwargs.get("top_p", self.config.top_p),
|
|
}
|
|
response = await self.client.chat.completions.create(**params)
|
|
return response.choices[0].message
|
|
|
|
|
|
def create_openai_client(
|
|
api_key: Optional[str] = None,
|
|
**kwargs: Any,
|
|
) -> OpenAIClient:
|
|
"""
|
|
便捷工厂:创建 OpenAIClient。
|
|
|
|
读取环境变量:
|
|
- OPENAI_API_KEY(必填)
|
|
- OPENAI_BASE_URL(可选)
|
|
- OPENAI_MODEL / OPENAI_CHAT_MODEL(可选)
|
|
"""
|
|
if api_key is None:
|
|
api_key = (os.getenv("OPENAI_API_KEY") or "").strip()
|
|
if not api_key:
|
|
raise ValueError("未提供 api_key 且环境变量 OPENAI_API_KEY 未设置。")
|
|
|
|
base_url = (kwargs.pop("base_url", None) or os.getenv("OPENAI_BASE_URL") or "").strip()
|
|
model_name = (
|
|
kwargs.pop("model_name", None)
|
|
or os.getenv("OPENAI_MODEL")
|
|
or os.getenv("OPENAI_CHAT_MODEL")
|
|
or ""
|
|
).strip()
|
|
|
|
cfg_kwargs: Dict[str, Any] = dict(kwargs)
|
|
if base_url:
|
|
cfg_kwargs["base_url"] = base_url
|
|
if model_name:
|
|
cfg_kwargs["model_name"] = model_name
|
|
|
|
config = OpenAIConfig(api_key=api_key, **cfg_kwargs)
|
|
return OpenAIClient(config)
|
|
|