Enhance SQL generation and streaming capabilities in the API server. Introduce optional parameters for streaming throttle and SQL stream granularity in NLChatRequest. Implement new functions for iterating SQL generation content pieces and adjusting streaming behavior based on user-defined settings. Update prompts for few-shot SQL adaptation and improve logging for SQL generation processes. Refactor orchestrator methods to support streaming responses and integrate few-shot SQL conditions. Update impact analysis documentation to reflect these changes.
This commit is contained in:
@@ -10,7 +10,7 @@ import json
|
||||
import logging
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional
|
||||
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]
|
||||
@@ -119,6 +119,69 @@ class OpenAIClient:
|
||||
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]:
|
||||
|
||||
Reference in New Issue
Block a user