""" 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, Iterator, 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_stream( self, messages: List[Dict[str, str]], **kwargs: Any, ) -> Iterator[str]: """流式聊天:按 completion 增量产出文本片段(与 chat 同参,固定 stream=True)。""" 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 ), "stream": True, "timeout": kwargs.get("timeout", self.config.timeout), } 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: stream = self.client.chat.completions.create(**params) for chunk in stream: if not chunk.choices: continue delta = chunk.choices[0].delta if delta and getattr(delta, "content", None): yield delta.content except Exception as e: 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: stream = self.client.chat.completions.create(**params) for chunk in stream: if not chunk.choices: continue delta = chunk.choices[0].delta if delta and getattr(delta, "content", None): yield delta.content return 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)