Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,551 @@
|
||||
"""
|
||||
Ollama 本地模型提供商适配器
|
||||
"""
|
||||
import logging
|
||||
from typing import AsyncGenerator, Dict, List, Any
|
||||
|
||||
from .base import BaseLLMProvider, LLMConfig, LLMMessage, LLMResponse, LLMStreamChunk, ToolCall
|
||||
from .registry import ProviderRegistry
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@ProviderRegistry.register
|
||||
class OllamaProvider(BaseLLMProvider):
|
||||
"""
|
||||
Ollama 本地模型提供商适配器
|
||||
|
||||
支持本地部署的开源模型
|
||||
"""
|
||||
|
||||
provider_type = 'ollama'
|
||||
provider_name = 'Ollama (本地)'
|
||||
supported_model_types = ['chat', 'embedding']
|
||||
|
||||
DEFAULT_API_BASE = 'http://localhost:11434/v1'
|
||||
|
||||
def __init__(self, api_key: str = '', api_base: str = '', **kwargs):
|
||||
super().__init__(api_key, api_base, **kwargs)
|
||||
# Ollama 使用 ollama_host 参数
|
||||
ollama_host = kwargs.get('ollama_host', '')
|
||||
if ollama_host:
|
||||
self.api_base = f'{ollama_host.rstrip("/")}/v1'
|
||||
else:
|
||||
self.api_base = api_base or self.DEFAULT_API_BASE
|
||||
self._client = None
|
||||
self._async_client = None
|
||||
|
||||
def _get_client(self):
|
||||
"""获取同步客户端"""
|
||||
if self._client is None:
|
||||
try:
|
||||
from openai import OpenAI
|
||||
self._client = OpenAI(
|
||||
api_key='ollama', # Ollama 不需要 API Key
|
||||
base_url=self.api_base,
|
||||
)
|
||||
except ImportError:
|
||||
raise ImportError('请安装 openai 库: pip install openai')
|
||||
return self._client
|
||||
|
||||
def _get_async_client(self):
|
||||
"""获取异步客户端"""
|
||||
if self._async_client is None:
|
||||
try:
|
||||
from openai import AsyncOpenAI
|
||||
self._async_client = AsyncOpenAI(
|
||||
api_key='ollama', # Ollama 不需要 API Key
|
||||
base_url=self.api_base,
|
||||
)
|
||||
except ImportError:
|
||||
raise ImportError('请安装 openai 库: pip install openai')
|
||||
return self._async_client
|
||||
|
||||
def validate_config(self) -> bool:
|
||||
"""Ollama 不需要 API Key"""
|
||||
return True
|
||||
|
||||
def chat(
|
||||
self,
|
||||
messages: List[LLMMessage],
|
||||
config: LLMConfig,
|
||||
) -> LLMResponse:
|
||||
"""同步对话"""
|
||||
client = self._get_client()
|
||||
|
||||
# 构建请求参数
|
||||
kwargs = {
|
||||
'model': config.model,
|
||||
'messages': [m.to_dict() for m in messages],
|
||||
'temperature': config.temperature,
|
||||
'top_p': config.top_p,
|
||||
'max_tokens': config.max_tokens,
|
||||
'stop': config.stop,
|
||||
}
|
||||
|
||||
# 添加 Function Calling 参数(Ollama 部分模型支持)
|
||||
if config.tools:
|
||||
kwargs['tools'] = [t.to_openai_format() for t in config.tools]
|
||||
if config.tool_choice == 'required':
|
||||
kwargs['tool_choice'] = 'required'
|
||||
elif config.tool_choice == 'none':
|
||||
kwargs['tool_choice'] = 'none'
|
||||
elif config.tool_choice != 'auto':
|
||||
kwargs['tool_choice'] = {'type': 'function', 'function': {'name': config.tool_choice}}
|
||||
else:
|
||||
kwargs['tool_choice'] = 'auto'
|
||||
|
||||
response = client.chat.completions.create(**kwargs)
|
||||
|
||||
choice = response.choices[0]
|
||||
usage = response.usage
|
||||
|
||||
# 解析工具调用
|
||||
tool_calls = None
|
||||
if choice.message.tool_calls:
|
||||
tool_calls = self._parse_tool_calls(choice.message.tool_calls)
|
||||
|
||||
return LLMResponse(
|
||||
content=choice.message.content or '',
|
||||
model=response.model,
|
||||
prompt_tokens=usage.prompt_tokens if usage else 0,
|
||||
completion_tokens=usage.completion_tokens if usage else 0,
|
||||
total_tokens=usage.total_tokens if usage else 0,
|
||||
finish_reason=choice.finish_reason or '',
|
||||
raw_response=response.model_dump(),
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
|
||||
def _parse_tool_calls(self, tool_calls) -> List[ToolCall]:
|
||||
"""解析工具调用"""
|
||||
import json
|
||||
result = []
|
||||
for tc in tool_calls:
|
||||
try:
|
||||
arguments = json.loads(tc.function.arguments) if tc.function.arguments else {}
|
||||
except json.JSONDecodeError:
|
||||
arguments = {'raw': tc.function.arguments}
|
||||
|
||||
result.append(ToolCall(
|
||||
id=tc.id or f'call_{len(result)}',
|
||||
name=tc.function.name,
|
||||
arguments=arguments,
|
||||
))
|
||||
return result
|
||||
|
||||
async def chat_async(
|
||||
self,
|
||||
messages: List[LLMMessage],
|
||||
config: LLMConfig,
|
||||
) -> LLMResponse:
|
||||
"""异步对话"""
|
||||
client = self._get_async_client()
|
||||
|
||||
# 构建请求参数
|
||||
kwargs = {
|
||||
'model': config.model,
|
||||
'messages': [m.to_dict() for m in messages],
|
||||
'temperature': config.temperature,
|
||||
'top_p': config.top_p,
|
||||
'max_tokens': config.max_tokens,
|
||||
'stop': config.stop,
|
||||
}
|
||||
|
||||
# 添加 Function Calling 参数
|
||||
if config.tools:
|
||||
kwargs['tools'] = [t.to_openai_format() for t in config.tools]
|
||||
if config.tool_choice == 'required':
|
||||
kwargs['tool_choice'] = 'required'
|
||||
elif config.tool_choice == 'none':
|
||||
kwargs['tool_choice'] = 'none'
|
||||
elif config.tool_choice != 'auto':
|
||||
kwargs['tool_choice'] = {'type': 'function', 'function': {'name': config.tool_choice}}
|
||||
else:
|
||||
kwargs['tool_choice'] = 'auto'
|
||||
|
||||
response = await client.chat.completions.create(**kwargs)
|
||||
|
||||
choice = response.choices[0]
|
||||
usage = response.usage
|
||||
|
||||
# 解析工具调用
|
||||
tool_calls = None
|
||||
if choice.message.tool_calls:
|
||||
tool_calls = self._parse_tool_calls(choice.message.tool_calls)
|
||||
|
||||
return LLMResponse(
|
||||
content=choice.message.content or '',
|
||||
model=response.model,
|
||||
prompt_tokens=usage.prompt_tokens if usage else 0,
|
||||
completion_tokens=usage.completion_tokens if usage else 0,
|
||||
total_tokens=usage.total_tokens if usage else 0,
|
||||
finish_reason=choice.finish_reason or '',
|
||||
raw_response=response.model_dump(),
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
|
||||
async def chat_stream(
|
||||
self,
|
||||
messages: List[LLMMessage],
|
||||
config: LLMConfig,
|
||||
) -> AsyncGenerator[LLMStreamChunk, None]:
|
||||
"""流式对话"""
|
||||
client = self._get_async_client()
|
||||
|
||||
# 构建请求参数
|
||||
kwargs = {
|
||||
'model': config.model,
|
||||
'messages': [m.to_dict() for m in messages],
|
||||
'temperature': config.temperature,
|
||||
'top_p': config.top_p,
|
||||
'max_tokens': config.max_tokens,
|
||||
'stop': config.stop,
|
||||
'stream': True,
|
||||
}
|
||||
|
||||
# 添加 Function Calling 参数
|
||||
if config.tools:
|
||||
kwargs['tools'] = [t.to_openai_format() for t in config.tools]
|
||||
if config.tool_choice == 'required':
|
||||
kwargs['tool_choice'] = 'required'
|
||||
elif config.tool_choice == 'none':
|
||||
kwargs['tool_choice'] = 'none'
|
||||
elif config.tool_choice != 'auto':
|
||||
kwargs['tool_choice'] = {'type': 'function', 'function': {'name': config.tool_choice}}
|
||||
else:
|
||||
kwargs['tool_choice'] = 'auto'
|
||||
|
||||
stream = await client.chat.completions.create(**kwargs)
|
||||
|
||||
# 用于累积工具调用
|
||||
tool_call_accumulator = {}
|
||||
|
||||
async for chunk in stream:
|
||||
if chunk.choices:
|
||||
choice = chunk.choices[0]
|
||||
delta = choice.delta
|
||||
|
||||
# 处理工具调用增量
|
||||
tool_call_delta = None
|
||||
if hasattr(delta, 'tool_calls') and delta.tool_calls:
|
||||
for tc_delta in delta.tool_calls:
|
||||
tc_index = tc_delta.index if hasattr(tc_delta, 'index') else 0
|
||||
if tc_index not in tool_call_accumulator:
|
||||
tool_call_accumulator[tc_index] = {
|
||||
'id': getattr(tc_delta, 'id', '') or f'call_{tc_index}',
|
||||
'name': '',
|
||||
'arguments': '',
|
||||
}
|
||||
if hasattr(tc_delta, 'id') and tc_delta.id:
|
||||
tool_call_accumulator[tc_index]['id'] = tc_delta.id
|
||||
if tc_delta.function:
|
||||
if tc_delta.function.name:
|
||||
tool_call_accumulator[tc_index]['name'] = tc_delta.function.name
|
||||
if tc_delta.function.arguments:
|
||||
tool_call_accumulator[tc_index]['arguments'] += tc_delta.function.arguments
|
||||
tool_call_delta = tool_call_accumulator[tc_index]
|
||||
|
||||
# 检查是否完成
|
||||
tool_calls = None
|
||||
if choice.finish_reason and tool_call_accumulator:
|
||||
tool_calls = self._parse_accumulated_tool_calls(tool_call_accumulator)
|
||||
|
||||
yield LLMStreamChunk(
|
||||
content=delta.content or '',
|
||||
is_finished=choice.finish_reason is not None,
|
||||
finish_reason=choice.finish_reason or '',
|
||||
tool_calls=tool_calls,
|
||||
tool_call_delta=tool_call_delta,
|
||||
)
|
||||
|
||||
def _parse_accumulated_tool_calls(self, accumulator: Dict) -> List[ToolCall]:
|
||||
"""解析累积的工具调用"""
|
||||
import json
|
||||
result = []
|
||||
for idx in sorted(accumulator.keys()):
|
||||
tc = accumulator[idx]
|
||||
try:
|
||||
arguments = json.loads(tc['arguments']) if tc['arguments'] else {}
|
||||
except json.JSONDecodeError:
|
||||
arguments = {'raw': tc['arguments']}
|
||||
|
||||
result.append(ToolCall(
|
||||
id=tc['id'],
|
||||
name=tc['name'],
|
||||
arguments=arguments,
|
||||
))
|
||||
return result
|
||||
|
||||
def chat_stream_sync(
|
||||
self,
|
||||
messages: List[LLMMessage],
|
||||
config: LLMConfig,
|
||||
):
|
||||
"""同步流式对话"""
|
||||
client = self._get_client()
|
||||
|
||||
# 构建请求参数
|
||||
kwargs = {
|
||||
'model': config.model,
|
||||
'messages': [m.to_dict() for m in messages],
|
||||
'temperature': config.temperature,
|
||||
'top_p': config.top_p,
|
||||
'max_tokens': config.max_tokens,
|
||||
'stop': config.stop,
|
||||
'stream': True,
|
||||
}
|
||||
|
||||
# 添加 Function Calling 参数
|
||||
if config.tools:
|
||||
kwargs['tools'] = [t.to_openai_format() for t in config.tools]
|
||||
if config.tool_choice == 'required':
|
||||
kwargs['tool_choice'] = 'required'
|
||||
elif config.tool_choice == 'none':
|
||||
kwargs['tool_choice'] = 'none'
|
||||
elif config.tool_choice != 'auto':
|
||||
kwargs['tool_choice'] = {'type': 'function', 'function': {'name': config.tool_choice}}
|
||||
else:
|
||||
kwargs['tool_choice'] = 'auto'
|
||||
|
||||
stream = client.chat.completions.create(**kwargs)
|
||||
|
||||
# 用于累积工具调用
|
||||
tool_call_accumulator = {}
|
||||
|
||||
for chunk in stream:
|
||||
if chunk.choices:
|
||||
choice = chunk.choices[0]
|
||||
delta = choice.delta
|
||||
|
||||
# 处理工具调用增量
|
||||
tool_call_delta = None
|
||||
if hasattr(delta, 'tool_calls') and delta.tool_calls:
|
||||
for tc_delta in delta.tool_calls:
|
||||
tc_index = tc_delta.index if hasattr(tc_delta, 'index') else 0
|
||||
if tc_index not in tool_call_accumulator:
|
||||
tool_call_accumulator[tc_index] = {
|
||||
'id': getattr(tc_delta, 'id', '') or f'call_{tc_index}',
|
||||
'name': '',
|
||||
'arguments': '',
|
||||
}
|
||||
if hasattr(tc_delta, 'id') and tc_delta.id:
|
||||
tool_call_accumulator[tc_index]['id'] = tc_delta.id
|
||||
if tc_delta.function:
|
||||
if tc_delta.function.name:
|
||||
tool_call_accumulator[tc_index]['name'] = tc_delta.function.name
|
||||
if tc_delta.function.arguments:
|
||||
tool_call_accumulator[tc_index]['arguments'] += tc_delta.function.arguments
|
||||
tool_call_delta = tool_call_accumulator[tc_index]
|
||||
|
||||
# 检查是否完成
|
||||
tool_calls = None
|
||||
if choice.finish_reason and tool_call_accumulator:
|
||||
tool_calls = self._parse_accumulated_tool_calls(tool_call_accumulator)
|
||||
|
||||
yield LLMStreamChunk(
|
||||
content=delta.content or '',
|
||||
is_finished=choice.finish_reason is not None,
|
||||
finish_reason=choice.finish_reason or '',
|
||||
tool_calls=tool_calls,
|
||||
tool_call_delta=tool_call_delta,
|
||||
)
|
||||
|
||||
_EMBEDDING_KEYWORDS = ('embed', 'bge', 'nomic-embed', 'mxbai-embed')
|
||||
_RERANK_KEYWORDS = ('rerank', 'reranker')
|
||||
|
||||
def get_available_models(self) -> List[Dict[str, Any]]:
|
||||
"""获取本地可用的模型列表"""
|
||||
try:
|
||||
import httpx
|
||||
base_url = self.api_base.replace('/v1', '')
|
||||
response = httpx.get(f'{base_url}/api/tags', timeout=10)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
models = []
|
||||
for model in data.get('models', []):
|
||||
models.append({
|
||||
'model_name': model.get('name', ''),
|
||||
'display_name': model.get('name', ''),
|
||||
'size': model.get('size', 0),
|
||||
'modified_at': model.get('modified_at', ''),
|
||||
})
|
||||
return models
|
||||
except Exception as e:
|
||||
logger.warning(f'获取 Ollama 模型列表失败: {e}')
|
||||
return []
|
||||
|
||||
async def fetch_models_from_api(self) -> List[Dict[str, Any]]:
|
||||
"""通过 /api/tags 端点拉取 Ollama 本地已安装的模型"""
|
||||
import httpx
|
||||
|
||||
base_url = (self.api_base or 'http://localhost:11434').rstrip('/')
|
||||
# Ollama 的 /api/tags 在根路径,去掉 /v1
|
||||
if base_url.endswith('/v1'):
|
||||
base_url = base_url[:-3]
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10) as client:
|
||||
resp = await client.get(f'{base_url}/api/tags')
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
|
||||
results = []
|
||||
for item in data.get('models', []):
|
||||
name = item.get('name', '')
|
||||
if not name:
|
||||
continue
|
||||
|
||||
name_lower = name.lower()
|
||||
if any(kw in name_lower for kw in self._RERANK_KEYWORDS):
|
||||
model_type = 'rerank'
|
||||
elif any(kw in name_lower for kw in self._EMBEDDING_KEYWORDS):
|
||||
model_type = 'embedding'
|
||||
else:
|
||||
model_type = 'chat'
|
||||
|
||||
# 从 size 推算显示名称
|
||||
size_gb = round(item.get('size', 0) / (1024 ** 3), 1)
|
||||
display = f'{name} ({size_gb}GB)' if size_gb > 0 else name
|
||||
|
||||
results.append({
|
||||
'model_name': name,
|
||||
'display_name': display,
|
||||
'model_type': model_type,
|
||||
'max_tokens': 4096,
|
||||
'context_window': 4096,
|
||||
'supports_vision': False,
|
||||
'supports_function_call': model_type == 'chat',
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
})
|
||||
|
||||
type_order = {'chat': 0, 'embedding': 1, 'rerank': 2}
|
||||
results.sort(key=lambda m: (type_order.get(m['model_type'], 9), m['model_name']))
|
||||
logger.info(f'从 Ollama ({base_url}) 拉取到 {len(results)} 个本地模型')
|
||||
return results
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f'在线拉取 Ollama 模型列表失败 ({base_url}): {e}')
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def get_default_models(cls) -> List[Dict[str, Any]]:
|
||||
"""获取默认模型列表(常用的开源模型)"""
|
||||
return [
|
||||
# ---- Chat 模型 ----
|
||||
{
|
||||
'model_name': 'qwen3:32b',
|
||||
'display_name': 'Qwen3 32B',
|
||||
'model_type': 'chat',
|
||||
'max_tokens': 8192,
|
||||
'context_window': 131072,
|
||||
'supports_vision': False,
|
||||
'supports_function_call': True,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
{
|
||||
'model_name': 'qwen3:8b',
|
||||
'display_name': 'Qwen3 8B',
|
||||
'model_type': 'chat',
|
||||
'max_tokens': 8192,
|
||||
'context_window': 131072,
|
||||
'supports_vision': False,
|
||||
'supports_function_call': True,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
{
|
||||
'model_name': 'llama3.3',
|
||||
'display_name': 'Llama 3.3 70B',
|
||||
'model_type': 'chat',
|
||||
'max_tokens': 4096,
|
||||
'context_window': 128000,
|
||||
'supports_vision': False,
|
||||
'supports_function_call': True,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
{
|
||||
'model_name': 'llama3.2',
|
||||
'display_name': 'Llama 3.2',
|
||||
'model_type': 'chat',
|
||||
'max_tokens': 4096,
|
||||
'context_window': 128000,
|
||||
'supports_vision': False,
|
||||
'supports_function_call': True,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
{
|
||||
'model_name': 'deepseek-r1',
|
||||
'display_name': 'DeepSeek R1',
|
||||
'model_type': 'chat',
|
||||
'max_tokens': 4096,
|
||||
'context_window': 64000,
|
||||
'supports_vision': False,
|
||||
'supports_function_call': False,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
{
|
||||
'model_name': 'mistral',
|
||||
'display_name': 'Mistral',
|
||||
'model_type': 'chat',
|
||||
'max_tokens': 4096,
|
||||
'context_window': 32000,
|
||||
'supports_vision': False,
|
||||
'supports_function_call': True,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
{
|
||||
'model_name': 'gemma3',
|
||||
'display_name': 'Gemma 3',
|
||||
'model_type': 'chat',
|
||||
'max_tokens': 8192,
|
||||
'context_window': 128000,
|
||||
'supports_vision': True,
|
||||
'supports_function_call': True,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
# ---- Embedding 模型 ----
|
||||
{
|
||||
'model_name': 'nomic-embed-text',
|
||||
'display_name': 'Nomic Embed Text',
|
||||
'model_type': 'embedding',
|
||||
'max_tokens': 8192,
|
||||
'context_window': 8192,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
{
|
||||
'model_name': 'bge-m3',
|
||||
'display_name': 'BGE-M3',
|
||||
'model_type': 'embedding',
|
||||
'max_tokens': 8192,
|
||||
'context_window': 8192,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
{
|
||||
'model_name': 'mxbai-embed-large',
|
||||
'display_name': 'MxBai Embed Large',
|
||||
'model_type': 'embedding',
|
||||
'max_tokens': 512,
|
||||
'context_window': 512,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
# ---- Rerank 模型 ----
|
||||
{
|
||||
'model_name': 'bge-reranker-v2-m3',
|
||||
'display_name': 'BGE Reranker V2 M3',
|
||||
'model_type': 'rerank',
|
||||
'max_tokens': 8192,
|
||||
'context_window': 8192,
|
||||
'input_price': 0,
|
||||
'output_price': 0,
|
||||
},
|
||||
]
|
||||
Reference in New Issue
Block a user