Files
ai-g3sb-backman2.0/backend/llm/openai_client.py
T

341 lines
12 KiB
Python
Raw Normal View History

"""
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)