Add chatId in ws connection
This commit is contained in:
@@ -126,6 +126,7 @@ class AgentConfig:
|
||||
system_prompt: str = "You are a helpful, friendly voice assistant."
|
||||
greeting: str | None = None
|
||||
greeting_mode: str = "generated"
|
||||
fastgpt_reconnect_greeting: str = "欢迎回来继续对话"
|
||||
response_state: ResponseStateConfig = field(default_factory=ResponseStateConfig)
|
||||
|
||||
|
||||
@@ -134,7 +135,7 @@ class LLMConfig:
|
||||
"""LLM backend selection via ``provider``.
|
||||
|
||||
Set ``provider`` to ``"openai"`` (alias ``"llm"``) for OpenAI-compatible chat
|
||||
completions, or ``"fastgpt"`` for FastGPT server-side memory via ``chat_id``.
|
||||
completions, or ``"fastgpt"`` for FastGPT server-side memory via runtime ``chatId``.
|
||||
"""
|
||||
|
||||
provider: str = "openai"
|
||||
|
||||
@@ -271,6 +271,39 @@ class FastGPTLLMService(LLMService):
|
||||
logger.warning(f"FastGPT chat init error: {exc}")
|
||||
return None
|
||||
|
||||
async def has_chat_history(self) -> bool:
|
||||
"""Return whether FastGPT has persisted records for this chatId."""
|
||||
if not self._app_id:
|
||||
return False
|
||||
|
||||
try:
|
||||
response = await self._client.get_chat_records(
|
||||
appId=self._app_id,
|
||||
chatId=self._chat_id,
|
||||
offset=0,
|
||||
pageSize=1,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
records = data.get("data", {}).get("list", [])
|
||||
return isinstance(records, list) and bool(records)
|
||||
except FastGPTError as exc:
|
||||
logger.warning(f"FastGPT chat records failed: {exc}")
|
||||
except httpx.HTTPError as exc:
|
||||
logger.warning(f"FastGPT chat records HTTP error: {exc}")
|
||||
except Exception as exc:
|
||||
logger.warning(f"FastGPT chat records error: {exc}")
|
||||
return False
|
||||
|
||||
async def fetch_session_greeting_text(self, reconnect_greeting: str) -> str | None:
|
||||
"""Use opener for a new chatId and a fixed greeting for reconnects."""
|
||||
if await self.has_chat_history():
|
||||
logger.info(f"FastGPT chatId={self._chat_id} has history; using reconnect greeting")
|
||||
return reconnect_greeting.strip() or None
|
||||
|
||||
logger.info(f"FastGPT chatId={self._chat_id} has no history; using app opener")
|
||||
return await self.fetch_welcome_text()
|
||||
|
||||
async def _close_active_response(self) -> None:
|
||||
response = self._active_response
|
||||
self._active_response = None
|
||||
|
||||
@@ -47,6 +47,18 @@ from .transcript_stream import ProductTranscriptStreamProcessor
|
||||
from .turn_start import InterruptionGateUserTurnStartStrategy
|
||||
|
||||
|
||||
def _chat_id_from_websocket(websocket) -> str | None:
|
||||
query_params = getattr(websocket, "query_params", None)
|
||||
if not query_params:
|
||||
return None
|
||||
|
||||
for name in ("chatId", "chat_id"):
|
||||
value = query_params.get(name)
|
||||
if isinstance(value, str) and value.strip():
|
||||
return value.strip()
|
||||
return None
|
||||
|
||||
|
||||
async def run_voice_pipeline(websocket, config: EngineConfig) -> None:
|
||||
await run_pipeline_with_serializer(
|
||||
websocket,
|
||||
@@ -93,7 +105,7 @@ async def run_pipeline_with_serializer(
|
||||
stt = create_stt_service(config.services.stt, config.audio)
|
||||
|
||||
llm_config = config.services.llm
|
||||
chat_id = llm_config.chat_id or f"voice_{uuid.uuid4().hex[:16]}"
|
||||
chat_id = _chat_id_from_websocket(websocket) or f"voice_{uuid.uuid4().hex[:16]}"
|
||||
llm = create_llm_service(
|
||||
llm_config,
|
||||
chat_id=chat_id,
|
||||
@@ -200,7 +212,9 @@ async def run_pipeline_with_serializer(
|
||||
await task.queue_frames([TTSSpeakFrame(config.agent.greeting)])
|
||||
elif config.agent.greeting_mode == "fastgpt_opener":
|
||||
if isinstance(llm, FastGPTLLMService):
|
||||
welcome = await llm.fetch_welcome_text()
|
||||
welcome = await llm.fetch_session_greeting_text(
|
||||
config.agent.fastgpt_reconnect_greeting
|
||||
)
|
||||
if welcome:
|
||||
await task.queue_frames([TTSSpeakFrame(welcome)])
|
||||
else:
|
||||
|
||||
@@ -103,10 +103,12 @@ class ProductWebsocketSerializer(FrameSerializer):
|
||||
|
||||
message_type = message.get("type")
|
||||
if message_type == "session.start":
|
||||
chat_id = message.get("chatId") or message.get("chat_id")
|
||||
return InputTransportMessageFrame(
|
||||
message={
|
||||
"type": "session.started",
|
||||
"protocol": self.protocol,
|
||||
"chatId": chat_id if isinstance(chat_id, str) else None,
|
||||
"audio": {
|
||||
"encoding": "pcm_s16le",
|
||||
"sample_rate": self._sample_rate,
|
||||
|
||||
@@ -61,7 +61,7 @@ def create_llm_service(
|
||||
return FastGPTLLMService(
|
||||
api_key=config.api_key,
|
||||
base_url=config.base_url or "http://localhost:3000",
|
||||
chat_id=chat_id or config.chat_id,
|
||||
chat_id=chat_id,
|
||||
app_id=config.app_id,
|
||||
greeting_prompt=greeting_prompt,
|
||||
timeout=config.timeout_sec,
|
||||
|
||||
Reference in New Issue
Block a user