Files
Nikita Tyukalov, ASUS, Linux 2bc979c9b4 Initial commit
2026-07-19 22:02:06 +03:00

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