Initial commit
This commit is contained in:
150
ai.py
Normal file
150
ai.py
Normal file
@@ -0,0 +1,150 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user