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