319 lines
10 KiB
Python
319 lines
10 KiB
Python
"""This module implements matrix bot"""
|
|
|
|
import asyncio
|
|
import json
|
|
import traceback
|
|
import time
|
|
import html
|
|
from datetime import datetime
|
|
|
|
from nio import ClientConfig, AsyncClient
|
|
from nio import LoginResponse, SyncResponse, RoomMessagesResponse
|
|
from nio import MatrixRoom, RoomMessageText
|
|
from nio import RoomGetStateEventResponse, RoomSendResponse
|
|
|
|
import config
|
|
import ai
|
|
|
|
#
|
|
# Data
|
|
#
|
|
_client: AsyncClient = None
|
|
_client_task: asyncio.Task = None
|
|
_next_batch: str = None
|
|
_message_tasks = set()
|
|
|
|
#
|
|
# Helpers
|
|
#
|
|
def _unix_time_to_readable_date(timestamp: int | float) -> str:
|
|
"""Convert UNIX time to readable date"""
|
|
dt = datetime.fromtimestamp(timestamp)
|
|
return dt.strftime("%d.%m.%Y, %H:%M")
|
|
|
|
async def _matrix_get_room_topic(room_id: str) -> str | None:
|
|
"""Returns room topic as string"""
|
|
try:
|
|
response = await _client.room_get_state_event(room_id, "m.room.topic")
|
|
if not isinstance(response, RoomGetStateEventResponse):
|
|
print("[!] Failed to get room topic: {response}")
|
|
return None
|
|
return response.content["m.topic"]["m.text"][0]["body"].replace("<br>", "\n").replace("</br>", "\n")
|
|
except:
|
|
return None
|
|
|
|
async def _matrix_send_text_to_room(room_id: str, text: str) -> str | None:
|
|
try:
|
|
response = await _client.room_send(
|
|
room_id=room_id,
|
|
message_type="m.room.message",
|
|
content={
|
|
"msgtype": "m.text",
|
|
"body": str(text)
|
|
}
|
|
)
|
|
if not isinstance(response, RoomSendResponse):
|
|
return None
|
|
return response.event_id
|
|
except:
|
|
traceback.print_exc()
|
|
return None
|
|
|
|
async def _matrix_get_last_messages(room_id: str, n: int = 200) -> list[tuple[str, str]]:
|
|
response = await _client.room_messages(
|
|
room_id=room_id,
|
|
start=_next_batch,
|
|
direction="b",
|
|
limit=n
|
|
)
|
|
if not isinstance(response, RoomMessagesResponse):
|
|
return []
|
|
messages = []
|
|
for e in response.chunk:
|
|
if not isinstance(e, RoomMessageText):
|
|
continue
|
|
if not e.body:
|
|
continue
|
|
if e.body.startswith("IGNORE"):
|
|
continue
|
|
messages.append((e.sender, e.body, _unix_time_to_readable_date(e.server_timestamp / 1000)))
|
|
return list(reversed(messages))
|
|
|
|
#
|
|
# Private
|
|
#
|
|
def _load_next_batch() -> str | None:
|
|
"""Read state from disk"""
|
|
try:
|
|
with open(config.MATRIX_STATE_JSON, "r") as f:
|
|
j = json.load(f)
|
|
if j["access_token"] != _client.access_token:
|
|
raise RuntimeError()
|
|
return j["next_batch"]
|
|
except:
|
|
return None
|
|
|
|
def _dump_next_batch(next_batch: str):
|
|
try:
|
|
j = {
|
|
"access_token": _client.access_token,
|
|
"next_batch": next_batch
|
|
}
|
|
with open(config.MATRIX_STATE_JSON, "w") as f:
|
|
json.dump(j, f)
|
|
return True
|
|
except:
|
|
return False
|
|
|
|
def _get_login_data() -> dict:
|
|
"""Returns data to use during login. This function decides whether to login or use existing access token."""
|
|
try:
|
|
with open(config.MATRIX_SESSION_JSON, "r") as f:
|
|
existing_data = json.load(f)
|
|
# check if required fields are present
|
|
if not isinstance(existing_data["homeserver"], str):
|
|
raise RuntimeError()
|
|
if not isinstance(existing_data["username"], str):
|
|
raise RuntimeError()
|
|
if not isinstance(existing_data["device_id"], str):
|
|
raise RuntimeError()
|
|
if not isinstance(existing_data["access_token"], str):
|
|
raise RuntimeError()
|
|
# check if data has changed
|
|
if existing_data["homeserver"] != config.MATRIX_HOMESERVER:
|
|
raise RuntimeError()
|
|
if existing_data["username"] != config.MATRIX_USERNAME:
|
|
raise RuntimeError()
|
|
return existing_data
|
|
except:
|
|
traceback.print_exc()
|
|
return {
|
|
"homeserver": config.MATRIX_HOMESERVER,
|
|
"username": config.MATRIX_USERNAME
|
|
}
|
|
|
|
def _write_login_data(login_response: LoginResponse) -> bool:
|
|
"""Writes login data to disk. It will be used on the next start."""
|
|
try:
|
|
j = {
|
|
"homeserver": config.MATRIX_HOMESERVER,
|
|
"username": login_response.user_id,
|
|
"device_id": login_response.device_id,
|
|
"access_token": login_response.access_token
|
|
}
|
|
with open(config.MATRIX_SESSION_JSON, "w") as f:
|
|
json.dump(j, f, indent=4)
|
|
return True
|
|
except:
|
|
traceback.print_exc()
|
|
return False
|
|
|
|
async def _on_message_handle(room_id: str, model: str, messages: list, whitelist: list[str], blacklist: list[str]):
|
|
temp_id = await _matrix_send_text_to_room(room_id, "IGNORE Calling ollama... (btw send `CUT` to reset chat memory)")
|
|
try:
|
|
ollama_response = await ai.chat(
|
|
model,
|
|
messages,
|
|
whitelist,
|
|
blacklist
|
|
)
|
|
await _client.room_redact(
|
|
room_id=room_id,
|
|
event_id=temp_id
|
|
)
|
|
await _matrix_send_text_to_room(room_id, str(ollama_response))
|
|
except Exception as e:
|
|
traceback.print_exc()
|
|
await _matrix_send_text_to_room(room_id, f"IGNORE Failed to call ollama: {e}")
|
|
|
|
|
|
|
|
async def _on_message(room: MatrixRoom, event: RoomMessageText) -> None:
|
|
"""This event is called when there is a new message."""
|
|
# ignore message by ourselves
|
|
if event.sender == config.MATRIX_USERNAME:
|
|
return
|
|
# cut
|
|
if event.body.lower() == "cut":
|
|
await _matrix_send_text_to_room(
|
|
room.room_id,
|
|
"IGNORE Chat memory is erased"
|
|
)
|
|
return
|
|
room_topic = await _matrix_get_room_topic(room.room_id)
|
|
if not room_topic:
|
|
await _matrix_send_text_to_room(
|
|
room.room_id,
|
|
"Missing room topic"
|
|
)
|
|
return
|
|
# check if required fields are present
|
|
try:
|
|
j = json.loads(room_topic)
|
|
if "ollama" not in j:
|
|
raise RuntimeError("Missing `root[\"ollama\"]`")
|
|
if "model" not in j["ollama"]:
|
|
raise RuntimeError("Missing `root[\"ollama\"][\"model\"]`")
|
|
if "tools" not in j:
|
|
raise RuntimeError("Missing `root[\"tools\"]`")
|
|
if "whitelist" not in j["tools"]:
|
|
raise RuntimeError("Missing `root[\"tools\"][\"whitelist\"]`")
|
|
if "blacklist" not in j["tools"]:
|
|
raise RuntimeError("Missing `root[\"tools\"][\"blacklist\"]`")
|
|
if "system" not in j:
|
|
raise RuntimeError("Missing `root[\"system\"]`")
|
|
except Exception as e:
|
|
traceback.print_exc()
|
|
await _matrix_send_text_to_room(
|
|
room.room_id,
|
|
f"Exception has occured: {e}"
|
|
)
|
|
return
|
|
# last messages
|
|
last_messages = (await _matrix_get_last_messages(room.room_id))[-20:]
|
|
for i in range(len(last_messages) - 1, -1, -1):
|
|
if last_messages[i][1].lower() == "cut":
|
|
last_messages = last_messages[i+1:]
|
|
break
|
|
# message text with timestamp
|
|
timestamp = _unix_time_to_readable_date(event.server_timestamp / 1000)
|
|
message_text = f"[Сообщение от {timestamp}]\n\n{event.body}"
|
|
# build messages history for ollama
|
|
messages = []
|
|
messages.append(ai.Message(ai.Role.SYSTEM, j["system"]))
|
|
for m in last_messages:
|
|
role = ai.Role.SYSTEM if m[0] == config.MATRIX_USERNAME else ai.Role.USER
|
|
body = f"[Сообщение от {m[2]}]\n\n{m[1]}"
|
|
messages.append(ai.Message(role, body))
|
|
messages.append(ai.Message(ai.Role.USER, message_text))
|
|
# handle in background
|
|
t = asyncio.create_task(_on_message_handle(
|
|
room.room_id,
|
|
j["ollama"]["model"],
|
|
messages,
|
|
j["tools"]["whitelist"],
|
|
j["tools"]["blacklist"]
|
|
))
|
|
_message_tasks.add(t)
|
|
t.add_done_callback(_message_tasks.discard)
|
|
|
|
async def _sync_loop() -> None:
|
|
"""This loop syncs forever."""
|
|
global _next_batch
|
|
last_save_time = 0
|
|
while True:
|
|
try:
|
|
response = await _client.sync(timeout=30000, since=_next_batch)
|
|
if isinstance(response, SyncResponse):
|
|
_next_batch = response.next_batch
|
|
d = time.time() - last_save_time
|
|
if d >= config.MATRIX_NEXT_BATCH_DUMP_PERIOD:
|
|
_dump_next_batch(_next_batch)
|
|
last_save_time = time.time()
|
|
else:
|
|
print(f"[!] Sync error: {response}")
|
|
await asyncio.sleep(5)
|
|
except asyncio.CancelledError:
|
|
raise
|
|
except:
|
|
traceback.print_exc()
|
|
await asyncio.sleep(1)
|
|
|
|
#
|
|
# Public
|
|
#
|
|
async def start() -> bool:
|
|
"""Start the bot using data in config.py or stored data"""
|
|
global _client, _client_task, _next_batch
|
|
# do not continue if the client already exists
|
|
if _client is not None:
|
|
return False
|
|
# get data to use for login
|
|
login_data = _get_login_data()
|
|
if "access_token" not in login_data:
|
|
print(f"[I] Authenticating as {login_data['username']}...")
|
|
login_data["password"] = input(" Password: ")
|
|
# create the client
|
|
_client = AsyncClient(
|
|
config.MATRIX_HOMESERVER,
|
|
config.MATRIX_USERNAME
|
|
)
|
|
# add callbacks
|
|
_client.add_event_callback(_on_message, RoomMessageText)
|
|
# authenticate
|
|
if "password" in login_data:
|
|
login_response = await _client.login(login_data["password"])
|
|
if not isinstance(login_response, LoginResponse):
|
|
print(f"[!] Failed to authenticate: {login_response}")
|
|
_client = None
|
|
return False
|
|
_write_login_data(login_response)
|
|
print("[I] Authenticated using password")
|
|
# already authenticated
|
|
else:
|
|
_client.access_token = login_data["access_token"]
|
|
_client.user_id = login_data["username"]
|
|
print("[I] Authenticated using access token")
|
|
# work
|
|
_next_batch = _load_next_batch()
|
|
_client_task = asyncio.create_task(_sync_loop())
|
|
return True
|
|
|
|
async def stop() -> bool:
|
|
"""Stop the bot"""
|
|
global _client, _client_task
|
|
if _client is None:
|
|
return False
|
|
try:
|
|
_client_task.cancel()
|
|
await _client_task
|
|
except:
|
|
pass
|
|
_client_task = None
|
|
try:
|
|
await _client.close()
|
|
except:
|
|
traceback.print_exc()
|
|
_dump_next_batch(_next_batch)
|
|
return True
|