Files
ai-agent-admin/backend-fastapi/ai_platform/providers/base.py
T

369 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
LLM 提供商基类
"""
import logging
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Any, AsyncGenerator, Dict, List, Optional
logger = logging.getLogger(__name__)
@dataclass
class LLMMessage:
"""LLM 消息
支持多模态内容格式:
- 纯文本: content 为字符串
- 多模态: content 为列表 [{"type": "text", "text": "..."}, {"type": "image_url", "image_url": {"url": "..."}}]
"""
role: str # system, user, assistant, tool
content: Any # str 或 List[Dict] (多模态内容)
name: Optional[str] = None
# Function Calling 相关
tool_calls: Optional[List[Dict[str, Any]]] = None # assistant 消息中的工具调用
tool_call_id: Optional[str] = None # tool 消息中的工具调用 ID
def to_dict(self) -> Dict[str, Any]:
result = {'role': self.role, 'content': self.content}
if self.name:
result['name'] = self.name
if self.tool_calls:
result['tool_calls'] = self.tool_calls
if self.tool_call_id:
result['tool_call_id'] = self.tool_call_id
return result
@staticmethod
def create_multimodal_content(text: str, attachments: List[Dict] = None) -> Any:
"""
创建多模态内容
Args:
text: 文本内容
attachments: 附件列表 [{"type": "image/file", "url": "...", "mime_type": "...", "base64": "..."}]
Returns:
str 或 List[Dict] 格式的内容
Note:
对于图片类型,优先使用 base64 格式(因为 LLM 无法访问内部 URL)
base64 格式: data:<mime_type>;base64,<base64_content>
"""
if not attachments:
return text
content = []
# 添加文本
if text:
content.append({"type": "text", "text": text})
# 添加附件
for att in attachments:
att_type = att.get('type', 'file')
mime_type = att.get('mime_type', '')
url = att.get('url', '')
base64_content = att.get('base64', '')
if att_type == 'image' or mime_type.startswith('image/'):
# 图片类型
text_content = att.get('text_content', '')
if base64_content:
# 多模态模型:使用 OpenAI 多模态格式
# 使用 data URL 格式
image_url = f"data:{mime_type};base64,{base64_content}"
content.append({
"type": "image_url",
"image_url": {"url": image_url}
})
elif text_content:
# 非多模态模型:使用 OCR 识别的文本内容
file_name = att.get('name', 'unknown')
content.append({
"type": "text",
"text": f"\n{text_content}"
})
elif url:
# 回退到 URL(仅适用于公网可访问的 URL)
content.append({
"type": "image_url",
"image_url": {"url": url}
})
else:
# 其他文件类型 - 提取文本内容或显示文件信息
file_name = att.get('name', 'unknown')
text_content = att.get('text_content', '')
if text_content:
# 有提取的文本内容,将其添加到消息中
content.append({
"type": "text",
"text": f"\n--- 文件: {file_name} ---\n{text_content}\n--- 文件结束 ---"
})
else:
# 无法提取内容,只显示文件信息
file_size = att.get('size', 0)
size_str = f"{file_size / 1024:.1f}KB" if file_size < 1024 * 1024 else f"{file_size / 1024 / 1024:.1f}MB"
content.append({
"type": "text",
"text": f"\n[附件: {file_name} ({size_str}),该文件类型暂不支持内容提取]"
})
return content if len(content) > 1 else (content[0].get('text', '') if content else text)
@dataclass
class ToolCall:
"""工具调用"""
id: str # 工具调用 ID
name: str # 工具/函数名称
arguments: Dict[str, Any] # 参数
def to_dict(self) -> Dict[str, Any]:
import json
return {
'id': self.id,
'type': 'function',
'function': {
'name': self.name,
'arguments': json.dumps(self.arguments, ensure_ascii=False) if isinstance(self.arguments, dict) else self.arguments,
}
}
@dataclass
class LLMResponse:
"""LLM 响应"""
content: str
model: str = ''
prompt_tokens: int = 0
completion_tokens: int = 0
total_tokens: int = 0
finish_reason: str = ''
raw_response: Optional[Dict] = None
# Function Calling 相关
tool_calls: Optional[List[ToolCall]] = None # 工具调用列表
@property
def is_complete(self) -> bool:
return self.finish_reason in ('stop', 'end_turn', 'length')
@property
def has_tool_calls(self) -> bool:
"""是否包含工具调用"""
return bool(self.tool_calls)
@dataclass
class LLMStreamChunk:
"""LLM 流式响应块"""
content: str = ''
is_finished: bool = False
finish_reason: str = ''
prompt_tokens: int = 0
completion_tokens: int = 0
total_tokens: int = 0
# Function Calling 相关
tool_calls: Optional[List[ToolCall]] = None # 工具调用(流式时可能分多次返回)
tool_call_delta: Optional[Dict[str, Any]] = None # 工具调用增量
@dataclass
class ToolDefinition:
"""工具定义(用于 Function Calling"""
name: str # 工具名称
description: str # 工具描述
parameters: Dict[str, Any] # 参数 Schema (JSON Schema 格式)
def to_openai_format(self) -> Dict[str, Any]:
"""转换为 OpenAI 格式"""
return {
'type': 'function',
'function': {
'name': self.name,
'description': self.description,
'parameters': self.parameters,
}
}
def to_claude_format(self) -> Dict[str, Any]:
"""转换为 Claude 格式"""
return {
'name': self.name,
'description': self.description,
'input_schema': self.parameters,
}
@dataclass
class LLMConfig:
"""LLM 配置"""
model: str
temperature: float = 0.7
top_p: float = 1.0
max_tokens: int = 2048
stop: Optional[List[str]] = None
presence_penalty: float = 0.0
frequency_penalty: float = 0.0
extra_params: Dict[str, Any] = field(default_factory=dict)
# Function Calling 相关
tools: Optional[List[ToolDefinition]] = None # 工具定义列表
tool_choice: str = 'auto' # 工具选择策略: auto, none, required, 或具体工具名
class BaseLLMProvider(ABC):
"""
LLM 提供商基类
所有提供商适配器必须继承此类并实现抽象方法
"""
# 提供商类型标识
provider_type: str = ''
# 提供商显示名称
provider_name: str = ''
# 支持的模型类型
supported_model_types: List[str] = ['chat']
def __init__(
self,
api_key: str = '',
api_base: str = '',
**kwargs
):
self.api_key = api_key
self.api_base = api_base
self.extra_config = kwargs
@abstractmethod
def chat(
self,
messages: List[LLMMessage],
config: LLMConfig,
) -> LLMResponse:
"""
同步对话
Args:
messages: 消息列表
config: 配置
Returns:
LLMResponse
"""
pass
@abstractmethod
async def chat_async(
self,
messages: List[LLMMessage],
config: LLMConfig,
) -> LLMResponse:
"""
异步对话
Args:
messages: 消息列表
config: 配置
Returns:
LLMResponse
"""
pass
@abstractmethod
async def chat_stream(
self,
messages: List[LLMMessage],
config: LLMConfig,
) -> AsyncGenerator[LLMStreamChunk, None]:
"""
异步流式对话
Args:
messages: 消息列表
config: 配置
Yields:
LLMStreamChunk
"""
pass
def chat_stream_sync(
self,
messages: List[LLMMessage],
config: LLMConfig,
):
"""
同步流式对话(生成器)
Args:
messages: 消息列表
config: 配置
Yields:
LLMStreamChunk
"""
# 默认实现:调用同步 chat 并返回单个结果
response = self.chat(messages, config)
yield LLMStreamChunk(
content=response.content,
is_finished=True,
finish_reason=response.finish_reason,
prompt_tokens=response.prompt_tokens,
completion_tokens=response.completion_tokens,
total_tokens=response.total_tokens,
)
def validate_config(self) -> bool:
"""
验证配置是否有效
Returns:
是否有效
"""
return bool(self.api_key)
def get_available_models(self) -> List[Dict[str, Any]]:
"""
获取可用模型列表
Returns:
模型列表
"""
return []
async def fetch_models_from_api(self) -> List[Dict[str, Any]]:
"""
从提供商 API 在线拉取最新模型列表
子类可覆盖此方法实现具体的 API 调用逻辑。
默认返回空列表,表示该提供商不支持在线拉取。
Returns:
模型列表,格式与 get_default_models 一致
"""
return []
async def fetch_models_from_api_strict(self) -> List[Dict[str, Any]]:
"""
严格拉取模型列表,用于连接测试。
子类可在这里保留上游异常,让调用方给出明确错误原因。
"""
return await self.fetch_models_from_api()
@classmethod
def get_default_models(cls) -> List[Dict[str, Any]]:
"""
获取默认模型列表(用于初始化)
Returns:
默认模型列表
"""
return []