Enhance LLM integration by adding OpenAI client support and enabling dynamic routing between DeepSeek and OpenAI services. Update environment configuration to include LLM_SERVICE_CODE for service selection, and modify API server to accommodate new request parameters for language and model. Implement streaming response improvements for chat interactions, allowing for segmented SSE output. Update documentation and impact analysis to reflect these changes.

This commit is contained in:
陈辅元
2026-04-16 09:15:01 +08:00
parent 85ad31348e
commit df33657099
17 changed files with 958 additions and 64 deletions
+340
View File
@@ -0,0 +1,340 @@
"""
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)