151 lines
4.9 KiB
Python
151 lines
4.9 KiB
Python
"""This module interfaces with OpenWebUI using OpenAI-compatible /chat/completions API"""
|
|
|
|
import asyncio
|
|
import json
|
|
import traceback
|
|
import os
|
|
|
|
import importlib.util
|
|
import inspect
|
|
import sys
|
|
import re
|
|
from typing import Callable
|
|
|
|
from dataclasses import dataclass
|
|
from enum import Enum
|
|
|
|
from ollama import AsyncClient
|
|
|
|
import config
|
|
|
|
class Role(Enum):
|
|
SYSTEM = "system"
|
|
USER = "user"
|
|
ASSISTANT = "assistant"
|
|
|
|
@dataclass
|
|
class Message:
|
|
role: Role
|
|
content: str
|
|
|
|
#
|
|
# Data
|
|
#
|
|
_tools = None
|
|
|
|
#
|
|
# Private
|
|
#
|
|
def _load_tools() -> list[Callable]:
|
|
"""Returns a list of all callables in tools.*"""
|
|
tools = []
|
|
tools_path = os.path.abspath("./tools/")
|
|
for tool_path in os.listdir(tools_path):
|
|
if not tool_path.endswith(".py"):
|
|
continue
|
|
module_name = tool_path[:-3]
|
|
file_path = os.path.join(tools_path, tool_path)
|
|
try:
|
|
spec = importlib.util.spec_from_file_location(module_name, file_path)
|
|
module = importlib.util.module_from_spec(spec)
|
|
sys.modules[module_name] = module
|
|
spec.loader.exec_module(module)
|
|
|
|
for name, obj in inspect.getmembers(module):
|
|
if inspect.isfunction(obj) and obj.__module__ == module.__name__ and not name.startswith("_"):
|
|
tools.append(obj)
|
|
except:
|
|
traceback.print_exc()
|
|
continue
|
|
return tools
|
|
|
|
#
|
|
# Public
|
|
#
|
|
async def chat(model: str, messages: list[Message], tools_whitelist: list[str] = [], tools_blacklist: list[str] = [], **kwargs) -> str | None:
|
|
"""Chat with ollama"""
|
|
# load tools if they are not loaded
|
|
global _tools
|
|
if _tools is None:
|
|
_tools = _load_tools()
|
|
try:
|
|
client = AsyncClient(
|
|
host=config.OLLAMA_URL,
|
|
headers=config.OLLAMA_HEADERS
|
|
)
|
|
# get tools list allowed for this chat
|
|
allowed_tools = []
|
|
# add all whitelisted tools (if whitelist is enabled)
|
|
if tools_whitelist:
|
|
for tool in _tools:
|
|
for pattern in tools_whitelist:
|
|
if re.match(pattern, tool.__name__):
|
|
print(f"[I] Allowing tool {tool.__name__} (matched by `{pattern}`)")
|
|
allowed_tools.append(tool)
|
|
break
|
|
# add all tools if whitelist is missing
|
|
else:
|
|
allowed_tools = list(_tools)
|
|
# remove blacklisted tools if blacklist is present
|
|
if tools_blacklist:
|
|
for tool in list(allowed_tools):
|
|
for pattern in tools_blacklist:
|
|
if not pattern:
|
|
continue
|
|
if re.match(pattern, tool.__name__):
|
|
print(f"[I] Removing tool {tool.__name__} (matched by `{pattern}`)")
|
|
allowed_tools.remove(tool)
|
|
break
|
|
# convert messages to valid format
|
|
messages = list({"role": m.role.value, "content": m.content} for m in messages)
|
|
# execute until done
|
|
while True:
|
|
response = await client.chat(
|
|
model=model,
|
|
messages=messages,
|
|
tools=allowed_tools,
|
|
**kwargs
|
|
)
|
|
if not response.message.tool_calls:
|
|
return response.message.content
|
|
messages.append(response.message)
|
|
calls = response.message.tool_calls
|
|
for call in calls:
|
|
tool_name = call.function.name
|
|
try:
|
|
print(f"[I] Calling tool `{tool_name}`...")
|
|
args = call.function.arguments
|
|
for a in args:
|
|
print(f"{a} = {args[a]}")
|
|
for t in allowed_tools:
|
|
if t.__name__ != tool_name:
|
|
continue
|
|
if inspect.iscoroutinefunction(t):
|
|
result = await t(**args)
|
|
else:
|
|
result = t(**args)
|
|
messages.append({
|
|
"role": "tool",
|
|
"tool_name": tool_name,
|
|
"content": f"{result}"
|
|
})
|
|
break
|
|
else:
|
|
print(f"[!] Tool `{tool_name}` is not found")
|
|
messages.append({
|
|
"role": "tool",
|
|
"tool_name": tool_name,
|
|
"content": f"Tool `{tool_name}` does not exist"
|
|
})
|
|
except Exception as e:
|
|
traceback.print_exc()
|
|
messages.append({
|
|
"role": "tool",
|
|
"tool_name": tool_name,
|
|
"content": f"Exception occured: {e}"
|
|
})
|
|
return response.message.content
|
|
except:
|
|
traceback.print_exc()
|
|
return None
|