213 lines
8.7 KiB
Python
213 lines
8.7 KiB
Python
import json
|
||
import aiohttp
|
||
import asyncio
|
||
import time
|
||
from typing import List, Dict, AsyncGenerator
|
||
from interfaces.llm import LLM
|
||
from utils.logger import session_logger
|
||
from config import settings
|
||
from services.interrupt_handler import interrupt_handler
|
||
from implementations.base_service import BaseService
|
||
|
||
class BaseLLM(LLM, BaseService):
|
||
def __init__(self, api_key: str, base_url: str, service_name: str, **kwargs):
|
||
BaseService.__init__(self, "llm", kwargs.get("selected_role"))
|
||
self.session = None
|
||
self.api_key = api_key
|
||
self.base_url = base_url
|
||
self.connected = False
|
||
self.closed = False
|
||
self.response = None
|
||
self.device_id = None
|
||
self.session_id = None
|
||
self.service_name = service_name
|
||
|
||
async def connect(self):
|
||
try:
|
||
if self.session:
|
||
await self.session.close()
|
||
self.session = aiohttp.ClientSession()
|
||
self.connected = True
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"[{self.service_name}] 连接建立失败: {str(e)}",
|
||
)
|
||
self.connected = False
|
||
|
||
async def prepare_request(
|
||
self, transcript: str, history: List[Dict[str, str]], selected_role: dict
|
||
) -> Dict:
|
||
start_time = time.perf_counter()
|
||
if (
|
||
hasattr(self, "_system_message")
|
||
and self._system_message.get("content") == selected_role["content"]
|
||
):
|
||
messages = [self._system_message] # 复用缓存的系统消息
|
||
else:
|
||
messages = [{"role": "system", "content": selected_role["content"]}]
|
||
self._system_message = messages[0] # 更新缓存
|
||
recent_history = history[-settings.max_conversation_history :]
|
||
for exchange in recent_history:
|
||
messages.append({"role": "user", "content": exchange["user"]})
|
||
messages.append({"role": "assistant", "content": exchange["assistant"]})
|
||
messages.append({"role": "user", "content": transcript})
|
||
data = {
|
||
"model": self.model_id,
|
||
"messages": messages,
|
||
"stream": True,
|
||
}
|
||
end_time = time.perf_counter()
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"[{self.service_name}] 请求准备完成: {len(messages)}条消息, 耗时: {end_time - start_time:.4f}秒",
|
||
)
|
||
return data
|
||
|
||
async def generate_response_stream(
|
||
self, transcript: str, history: List[Dict[str, str]], selected_role: dict
|
||
) -> AsyncGenerator[str, None]:
|
||
if not self.connected or not self.session:
|
||
await self.connect()
|
||
if not self.connected:
|
||
yield None
|
||
return
|
||
self.device_id = getattr(self, "device_id", "unknown")
|
||
self.session_id = getattr(self, "session_id", "unknown")
|
||
session_key = (
|
||
(self.device_id, self.session_id)
|
||
if self.device_id != "unknown" and self.session_id != "unknown"
|
||
else None
|
||
)
|
||
try:
|
||
if session_key:
|
||
self.closed = interrupt_handler.is_interrupted(session_key)
|
||
data = await self.prepare_request(transcript, history, selected_role)
|
||
if session_key:
|
||
await interrupt_handler.set_interrupt_state(session_key, False)
|
||
self.closed = False
|
||
session_logger.info(
|
||
self.device_id, self.session_id, f"[{self.service_name}] 开始生成回复"
|
||
)
|
||
headers = {
|
||
"Authorization": f"Bearer {self.api_key}",
|
||
"Content-Type": "application/json",
|
||
}
|
||
try:
|
||
if not self.session or self.session.closed:
|
||
self.session = aiohttp.ClientSession()
|
||
self.connected = True
|
||
response = await self.session.post(
|
||
f"{self.base_url}/chat/completions", headers=headers, json=data
|
||
)
|
||
self.response = response
|
||
if response.status == 200:
|
||
async for line in response.content:
|
||
if session_key and interrupt_handler.is_interrupted(
|
||
session_key
|
||
):
|
||
self.closed = True
|
||
if self.closed:
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"[{self.service_name}] 生成被中断",
|
||
)
|
||
break
|
||
decoded_line = line.decode("utf-8").strip()
|
||
if not decoded_line:
|
||
continue
|
||
if decoded_line.startswith("data:"):
|
||
data_str = decoded_line[len("data:") :].strip()
|
||
if data_str == "[DONE]":
|
||
break
|
||
try:
|
||
data_json = json.loads(data_str)
|
||
except json.JSONDecodeError:
|
||
continue
|
||
choices = data_json.get("choices", [])
|
||
if choices:
|
||
delta = choices[0].get("delta", {})
|
||
content = delta.get("content")
|
||
if content:
|
||
if self.closed:
|
||
session_logger.info(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"[{self.service_name}] 生成被中断",
|
||
)
|
||
break
|
||
yield content
|
||
else:
|
||
error_text = await response.text()
|
||
session_logger.error(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"[{self.service_name}] 生成回复失败: {error_text}",
|
||
)
|
||
yield None
|
||
except asyncio.CancelledError:
|
||
session_logger.info(
|
||
self.device_id, self.session_id, f"[{self.service_name}] 请求被取消"
|
||
)
|
||
yield None
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"[{self.service_name}] API请求错误: {str(e)}",
|
||
)
|
||
yield None
|
||
except Exception as e:
|
||
session_logger.error(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"[{self.service_name}] 生成过程发生错误: {str(e)}",
|
||
)
|
||
yield None
|
||
finally:
|
||
if self.response:
|
||
await self.response.release()
|
||
self.response = None
|
||
|
||
async def close(self):
|
||
if (
|
||
hasattr(self, "device_id")
|
||
and hasattr(self, "session_id")
|
||
and self.device_id != "unknown"
|
||
and self.session_id != "unknown"
|
||
):
|
||
session_key = (self.device_id, self.session_id)
|
||
await interrupt_handler.set_interrupt_state(session_key, True)
|
||
self.closed = True # 本地标志也同步设置
|
||
if self.response:
|
||
try:
|
||
self.response.close() # 立即关闭响应,不等待release
|
||
await self.response.release()
|
||
except Exception as e:
|
||
if hasattr(self, "device_id") and hasattr(self, "session_id"):
|
||
session_logger.warning(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"[{self.service_name}] 关闭响应时出错: {e}",
|
||
)
|
||
self.response = None
|
||
if self.session:
|
||
try:
|
||
await self.session.close()
|
||
except Exception as e:
|
||
if hasattr(self, "device_id") and hasattr(self, "session_id"):
|
||
session_logger.warning(
|
||
self.device_id,
|
||
self.session_id,
|
||
f"[{self.service_name}] 关闭会话时出错: {e}",
|
||
)
|
||
self.session = None
|
||
self.connected = False
|
||
if hasattr(self, "device_id") and hasattr(self, "session_id"):
|
||
session_logger.info(
|
||
self.device_id, self.session_id, f"[{self.service_name}] 连接已关闭"
|
||
)
|