432 lines
14 KiB
Python
432 lines
14 KiB
Python
"""
|
|
对话服务
|
|
"""
|
|
import logging
|
|
import time
|
|
from typing import AsyncGenerator, Dict, List, Optional
|
|
|
|
from sqlalchemy import select, func
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from ai_platform.models import AIApp, Conversation, Message, LLMModel, LLMProvider
|
|
from utils.context import get_current_user_id_from_context
|
|
from .llm_service import LLMService
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class ChatService:
|
|
"""
|
|
对话服务
|
|
|
|
管理对话和消息
|
|
"""
|
|
|
|
def __init__(self, db: AsyncSession):
|
|
self._db = db
|
|
self.llm_service = LLMService(db)
|
|
|
|
async def create_conversation(
|
|
self,
|
|
app_id: str,
|
|
title: str = '',
|
|
user_id: Optional[str] = None,
|
|
) -> Conversation:
|
|
"""
|
|
创建对话
|
|
|
|
Args:
|
|
app_id: 应用 ID
|
|
title: 对话标题
|
|
|
|
Returns:
|
|
Conversation
|
|
"""
|
|
current_user_id = user_id or get_current_user_id_from_context()
|
|
if not current_user_id:
|
|
raise ValueError('未登录或登录已过期')
|
|
|
|
# 查询应用
|
|
result = await self._db.execute(
|
|
select(AIApp).where(AIApp.id == app_id, AIApp.is_deleted == False)
|
|
)
|
|
app = result.scalar_one_or_none()
|
|
if not app:
|
|
raise ValueError(f'应用不存在: {app_id}')
|
|
|
|
conversation = Conversation(
|
|
app_id=app_id,
|
|
user_id=current_user_id,
|
|
title=title or '新对话',
|
|
)
|
|
self._db.add(conversation)
|
|
|
|
# 更新应用统计
|
|
app.conversation_count = (app.conversation_count or 0) + 1
|
|
|
|
await self._db.commit()
|
|
await self._db.refresh(conversation)
|
|
|
|
return conversation
|
|
|
|
async def get_conversation(self, conversation_id: str) -> Optional[Conversation]:
|
|
"""获取对话"""
|
|
result = await self._db.execute(
|
|
select(Conversation).where(
|
|
Conversation.id == conversation_id,
|
|
Conversation.is_deleted == False
|
|
)
|
|
)
|
|
return result.scalar_one_or_none()
|
|
|
|
async def list_conversations(
|
|
self,
|
|
app_id: str,
|
|
page: int = 1,
|
|
page_size: int = 20,
|
|
) -> tuple:
|
|
"""
|
|
获取对话列表
|
|
|
|
Returns:
|
|
(conversations, total)
|
|
"""
|
|
query = select(Conversation).where(
|
|
Conversation.app_id == app_id,
|
|
Conversation.is_deleted == False
|
|
).order_by(Conversation.is_pinned.desc(), Conversation.sys_update_datetime.desc())
|
|
|
|
# 获取总数
|
|
count_result = await self._db.execute(
|
|
select(func.count()).select_from(query.subquery())
|
|
)
|
|
total = count_result.scalar() or 0
|
|
|
|
# 分页
|
|
offset = (page - 1) * page_size
|
|
query = query.offset(offset).limit(page_size)
|
|
|
|
result = await self._db.execute(query)
|
|
conversations = result.scalars().all()
|
|
|
|
return list(conversations), total
|
|
|
|
async def delete_conversation(self, conversation_id: str) -> bool:
|
|
"""删除对话"""
|
|
result = await self._db.execute(
|
|
select(Conversation).where(
|
|
Conversation.id == conversation_id,
|
|
Conversation.is_deleted == False
|
|
)
|
|
)
|
|
conversation = result.scalar_one_or_none()
|
|
if not conversation:
|
|
return False
|
|
|
|
conversation.is_deleted = True
|
|
await self._db.commit()
|
|
return True
|
|
|
|
async def get_messages(
|
|
self,
|
|
conversation_id: str,
|
|
limit: int = 50,
|
|
) -> List[Message]:
|
|
"""获取对话消息"""
|
|
result = await self._db.execute(
|
|
select(Message).where(
|
|
Message.conversation_id == conversation_id,
|
|
Message.is_deleted == False
|
|
).order_by(Message.sys_create_datetime).limit(limit)
|
|
)
|
|
return list(result.scalars().all())
|
|
|
|
async def send_message(
|
|
self,
|
|
conversation_id: str,
|
|
content: str,
|
|
) -> tuple:
|
|
"""
|
|
发送消息并获取 AI 回复
|
|
|
|
Args:
|
|
conversation_id: 对话 ID
|
|
content: 消息内容
|
|
|
|
Returns:
|
|
(user_message, assistant_message)
|
|
"""
|
|
conversation = await self.get_conversation(conversation_id)
|
|
if not conversation:
|
|
raise ValueError('对话不存在')
|
|
|
|
# 获取应用
|
|
app_result = await self._db.execute(
|
|
select(AIApp).where(AIApp.id == conversation.app_id)
|
|
)
|
|
app = app_result.scalar_one_or_none()
|
|
if not app:
|
|
raise ValueError('应用不存在')
|
|
|
|
# 获取模型
|
|
model = await self._get_effective_model(conversation, app)
|
|
if not model:
|
|
raise ValueError('未配置模型')
|
|
|
|
# 创建用户消息
|
|
user_message = Message(
|
|
conversation_id=conversation_id,
|
|
role='user',
|
|
content=content,
|
|
status='completed',
|
|
)
|
|
self._db.add(user_message)
|
|
await self._db.flush()
|
|
|
|
# 构建消息列表
|
|
messages = await self._build_messages(conversation, app, content)
|
|
|
|
# 创建助手消息(pending 状态)
|
|
assistant_message = Message(
|
|
conversation_id=conversation_id,
|
|
role='assistant',
|
|
content='',
|
|
status='pending',
|
|
model_name=model.model_name,
|
|
)
|
|
self._db.add(assistant_message)
|
|
await self._db.flush()
|
|
|
|
try:
|
|
start_time = time.time()
|
|
|
|
# 调用 LLM
|
|
response = await self.llm_service.chat_async(
|
|
model_id=str(model.id),
|
|
messages=messages,
|
|
temperature=app.temperature or 0.7,
|
|
max_tokens=app.max_tokens or 2048,
|
|
)
|
|
|
|
latency = int((time.time() - start_time) * 1000)
|
|
|
|
# 更新助手消息
|
|
assistant_message.content = response.content
|
|
assistant_message.status = 'completed'
|
|
assistant_message.prompt_tokens = response.prompt_tokens
|
|
assistant_message.completion_tokens = response.completion_tokens
|
|
assistant_message.total_tokens = response.total_tokens
|
|
assistant_message.latency = latency
|
|
|
|
# 更新对话统计
|
|
conversation.message_count = (conversation.message_count or 0) + 2
|
|
conversation.total_tokens = (conversation.total_tokens or 0) + response.total_tokens
|
|
|
|
# 更新应用统计
|
|
app.message_count = (app.message_count or 0) + 2
|
|
|
|
# 自动生成标题
|
|
if conversation.message_count == 2:
|
|
self._generate_title(conversation, content)
|
|
|
|
await self._db.commit()
|
|
|
|
return user_message, assistant_message
|
|
|
|
except Exception as e:
|
|
logger.exception(f'发送消息失败: {e}')
|
|
assistant_message.status = 'failed'
|
|
assistant_message.error_message = str(e)
|
|
await self._db.commit()
|
|
raise
|
|
|
|
async def send_message_stream(
|
|
self,
|
|
conversation_id: str,
|
|
content: str,
|
|
) -> AsyncGenerator[Dict, None]:
|
|
"""
|
|
流式发送消息
|
|
|
|
Yields:
|
|
{"type": "content", "content": "..."} 或
|
|
{"type": "done", "message": {...}}
|
|
"""
|
|
conversation = await self.get_conversation(conversation_id)
|
|
if not conversation:
|
|
raise ValueError('对话不存在')
|
|
|
|
# 获取应用
|
|
app_result = await self._db.execute(
|
|
select(AIApp).where(AIApp.id == conversation.app_id)
|
|
)
|
|
app = app_result.scalar_one_or_none()
|
|
if not app:
|
|
raise ValueError('应用不存在')
|
|
|
|
# 获取模型
|
|
model = await self._get_effective_model(conversation, app)
|
|
if not model:
|
|
raise ValueError('未配置模型')
|
|
|
|
# 创建用户消息
|
|
user_message = Message(
|
|
conversation_id=conversation_id,
|
|
role='user',
|
|
content=content,
|
|
status='completed',
|
|
)
|
|
self._db.add(user_message)
|
|
await self._db.flush()
|
|
|
|
# 构建消息列表
|
|
messages = await self._build_messages(conversation, app, content)
|
|
|
|
# 创建助手消息
|
|
assistant_message = Message(
|
|
conversation_id=conversation_id,
|
|
role='assistant',
|
|
content='',
|
|
status='pending',
|
|
model_name=model.model_name,
|
|
)
|
|
self._db.add(assistant_message)
|
|
await self._db.flush()
|
|
|
|
try:
|
|
start_time = time.time()
|
|
full_content = ''
|
|
total_tokens = 0
|
|
prompt_tokens = 0
|
|
completion_tokens = 0
|
|
|
|
# 使用异步流式方法
|
|
async for chunk in self.llm_service.chat_stream(
|
|
model_id=str(model.id),
|
|
messages=messages,
|
|
temperature=app.temperature or 0.7,
|
|
max_tokens=app.max_tokens or 2048,
|
|
):
|
|
if chunk.content:
|
|
full_content += chunk.content
|
|
yield {'type': 'content', 'content': chunk.content}
|
|
|
|
if chunk.is_finished:
|
|
prompt_tokens = chunk.prompt_tokens
|
|
completion_tokens = chunk.completion_tokens
|
|
total_tokens = chunk.total_tokens
|
|
|
|
latency = int((time.time() - start_time) * 1000)
|
|
|
|
# 更新助手消息
|
|
assistant_message.content = full_content
|
|
assistant_message.status = 'completed'
|
|
assistant_message.prompt_tokens = prompt_tokens
|
|
assistant_message.completion_tokens = completion_tokens
|
|
assistant_message.total_tokens = total_tokens
|
|
assistant_message.latency = latency
|
|
|
|
# 更新对话统计
|
|
conversation.message_count = (conversation.message_count or 0) + 2
|
|
conversation.total_tokens = (conversation.total_tokens or 0) + total_tokens
|
|
|
|
await self._db.commit()
|
|
|
|
yield {
|
|
'type': 'done',
|
|
'message': {
|
|
'id': str(assistant_message.id),
|
|
'content': full_content,
|
|
'tokens': total_tokens,
|
|
'latency': latency,
|
|
},
|
|
}
|
|
|
|
except Exception as e:
|
|
logger.exception(f'流式发送消息失败: {e}')
|
|
assistant_message.status = 'failed'
|
|
assistant_message.error_message = str(e)
|
|
await self._db.commit()
|
|
yield {'type': 'error', 'error': str(e)}
|
|
|
|
async def _get_effective_model(self, conversation: Conversation, app: AIApp) -> Optional[LLMModel]:
|
|
"""获取有效的模型"""
|
|
model_id = conversation.model_override_id or app.model_id
|
|
if model_id:
|
|
result = await self._db.execute(
|
|
select(LLMModel)
|
|
.join(LLMProvider, LLMProvider.id == LLMModel.provider_id)
|
|
.where(
|
|
LLMModel.id == model_id,
|
|
LLMModel.is_active == True,
|
|
LLMModel.is_deleted == False,
|
|
LLMProvider.is_active == True,
|
|
LLMProvider.is_deleted == False,
|
|
)
|
|
)
|
|
model = result.scalar_one_or_none()
|
|
if model:
|
|
return model
|
|
|
|
result = await self._db.execute(
|
|
select(LLMModel)
|
|
.join(LLMProvider, LLMProvider.id == LLMModel.provider_id)
|
|
.where(
|
|
LLMModel.is_active == True,
|
|
LLMModel.is_deleted == False,
|
|
LLMModel.model_type == "chat",
|
|
LLMProvider.is_active == True,
|
|
LLMProvider.is_deleted == False,
|
|
)
|
|
.order_by(LLMModel.sort.desc(), LLMModel.sys_create_datetime.desc())
|
|
)
|
|
return result.scalar_one_or_none()
|
|
|
|
async def _build_messages(
|
|
self,
|
|
conversation: Conversation,
|
|
app: AIApp,
|
|
user_content: str,
|
|
) -> List[Dict[str, str]]:
|
|
"""构建消息列表"""
|
|
messages = []
|
|
|
|
# 系统提示词
|
|
if app.system_prompt:
|
|
messages.append({
|
|
'role': 'system',
|
|
'content': app.system_prompt,
|
|
})
|
|
|
|
# 历史消息
|
|
result = await self._db.execute(
|
|
select(Message).where(
|
|
Message.conversation_id == conversation.id,
|
|
Message.is_deleted == False,
|
|
Message.status == 'completed'
|
|
).order_by(Message.sys_create_datetime.desc()).limit(20)
|
|
)
|
|
history = result.scalars().all()
|
|
|
|
for msg in reversed(list(history)):
|
|
messages.append({
|
|
'role': msg.role,
|
|
'content': msg.content,
|
|
})
|
|
|
|
# 当前用户消息
|
|
messages.append({
|
|
'role': 'user',
|
|
'content': user_content,
|
|
})
|
|
|
|
return messages
|
|
|
|
def _generate_title(self, conversation: Conversation, first_message: str):
|
|
"""自动生成对话标题"""
|
|
# 简单截取前 20 个字符作为标题
|
|
title = first_message[:20]
|
|
if len(first_message) > 20:
|
|
title += '...'
|
|
conversation.title = title
|