nest-matrix-bot/matrix_keycloak_bot.py
2026-07-23 13:32:12 -04:00

464 lines
17 KiB
Python
Raw Blame History

This file contains invisible Unicode characters

This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import asyncio
import os
import sys
import markdown
import json
import pickle
import logging
from keycloak import KeycloakAdmin
from keycloak.exceptions import KeycloakError
from mautrix.client import Client
from mautrix.crypto import OlmMachine
from mautrix.crypto.store import MemoryCryptoStore
from mautrix.client.state_store import MemoryStateStore
from mautrix.types import (
EventType, MessageType, TextMessageEventContent, Format,
Event, RoomID, Membership, RelationType, UserID
)
# ------------------------------------------------------------------------------
# Configuration
# ------------------------------------------------------------------------------
MATRIX_HOMESERVER = os.getenv("MATRIX_HOMESERVER", "https://matrix.mynest.love")
MATRIX_BOT_USER = os.getenv("MATRIX_BOT_USER", "@access-bot:mynest.love")
MATRIX_BOT_PASSWORD = os.getenv("MATRIX_BOT_PASSWORD", "SuperSecretBotPassword")
ADMIN_ROOM_ID = os.getenv("ADMIN_ROOM_ID", "!adminroomid:mynest.love")
ADMIN_MATRIX_IDS = set(os.getenv("ADMIN_MATRIX_IDS", "@you:mynest.love").split(","))
KEYCLOAK_URL = os.getenv("KEYCLOAK_URL", "https://keycloak.mynest.love/")
KEYCLOAK_REALM = os.getenv("KEYCLOAK_REALM", "master")
KEYCLOAK_CLIENT_ID = os.getenv("KEYCLOAK_CLIENT_ID", "nest-matrix-bot")
KEYCLOAK_CLIENT_SECRET = os.getenv("KEYCLOAK_CLIENT_SECRET", "your-client-secret")
SESSION_FILE = os.getenv("MATRIX_SESSION_FILE", "/app/store/bot_session.json")
CRYPTO_STORE_FILE = os.getenv("MATRIX_CRYPTO_STORE", "/app/store/crypto_store.pkl")
# In-memory tracking of active requests (Event ID -> Request Info)
pending_requests = {}
logging.basicConfig(level=logging.INFO)
log = logging.getLogger("keycloak_bot")
# ------------------------------------------------------------------------------
# Custom Pickle-Backed Crypto Store for Persistence
# ------------------------------------------------------------------------------
class PickleCryptoStore(MemoryCryptoStore):
def __init__(self, filepath: str):
super().__init__()
self.filepath = filepath
self.load()
def load(self):
if os.path.exists(self.filepath):
try:
with open(self.filepath, "rb") as f:
data = pickle.load(f)
self.accounts = data.get("accounts", {})
self.sessions = data.get("sessions", {})
self.inbound_group_sessions = data.get("inbound_group_sessions", {})
self.outbound_group_sessions = data.get("outbound_group_sessions", {})
self.devices = data.get("devices", {})
self.cross_signing_keys = data.get("cross_signing_keys", {})
log.info(f"Loaded crypto store from {self.filepath}")
except Exception as e:
log.warning(f"Failed to load pickle crypto store: {e}")
def save(self):
try:
data = {
"accounts": self.accounts,
"sessions": self.sessions,
"inbound_group_sessions": self.inbound_group_sessions,
"outbound_group_sessions": self.outbound_group_sessions,
"devices": self.devices,
"cross_signing_keys": self.cross_signing_keys,
}
directory = os.path.dirname(self.filepath)
if directory:
os.makedirs(directory, exist_ok=True)
with open(self.filepath, "wb") as f:
pickle.dump(data, f)
except Exception as e:
log.error(f"Failed to save pickle crypto store: {e}")
async def put_account(self, *args, **kwargs):
await super().put_account(*args, **kwargs)
self.save()
async def put_session(self, *args, **kwargs):
await super().put_session(*args, **kwargs)
self.save()
async def put_inbound_group_session(self, *args, **kwargs):
await super().put_inbound_group_session(*args, **kwargs)
self.save()
async def put_outbound_group_session(self, *args, **kwargs):
await super().put_outbound_group_session(*args, **kwargs)
self.save()
async def put_device(self, *args, **kwargs):
await super().put_device(*args, **kwargs)
self.save()
async def put_cross_signing_key(self, user_id, usage, key):
await super().put_cross_signing_key(user_id, usage, key)
self.save()
# ------------------------------------------------------------------------------
# Custom State Store supporting find_shared_rooms
# ------------------------------------------------------------------------------
class CustomMemoryStateStore(MemoryStateStore):
def __init__(self, bot_mxid: str):
super().__init__()
self.bot_mxid = bot_mxid
async def find_shared_rooms(self, user_id: UserID) -> list[RoomID]:
shared = []
states = getattr(self, "states", {})
for room_id in states.keys():
try:
rid = RoomID(room_id)
user_membership = await self.get_membership(rid, user_id)
bot_membership = await self.get_membership(rid, UserID(self.bot_mxid))
if user_membership == Membership.JOIN and bot_membership == Membership.JOIN:
shared.append(rid)
except Exception:
pass
return shared
# ------------------------------------------------------------------------------
# Keycloak Helper Functions
# ------------------------------------------------------------------------------
def get_keycloak_client() -> KeycloakAdmin:
"""Initializes KeycloakAdmin client using client credentials."""
return KeycloakAdmin(
server_url=KEYCLOAK_URL,
client_id=KEYCLOAK_CLIENT_ID,
client_secret_key=KEYCLOAK_CLIENT_SECRET,
realm_name=KEYCLOAK_REALM,
user_realm_name=KEYCLOAK_REALM,
grant_type="client_credentials",
verify=True,
)
def check_group_membership(username: str, group_name: str) -> tuple[bool, str]:
try:
kc = get_keycloak_client()
users = kc.get_users(query={"username": username, "exact": True})
if not users:
return False, f"⚠️ Keycloak user `{username}` not found."
user_id = users[0]["id"]
user_groups = kc.get_user_groups(user_id=user_id)
user_group_names = {g["name"].lower() for g in user_groups}
if group_name.lower() in user_group_names:
return True, f" You are already a member of the `{group_name}` group."
return False, ""
except KeycloakError as e:
return False, f"⚠️ Keycloak API error: {str(e)}"
except Exception as e:
return False, f"⚠️ Unexpected error: {str(e)}"
def add_user_to_kc_group(username: str, group_name: str) -> tuple[bool, str]:
try:
kc = get_keycloak_client()
users = kc.get_users(query={"username": username, "exact": True})
if not users:
return False, f"Keycloak user '{username}' not found."
user_id = users[0]["id"]
groups = kc.get_groups()
target_group = next((g for g in groups if g["name"].lower() == group_name.lower()), None)
if not target_group:
return False, f"Keycloak group '{group_name}' not found."
group_id = target_group["id"]
kc.group_user_add(user_id=user_id, group_id=group_id)
return True, f"Successfully added '{username}' to group '{target_group['name']}'."
except KeycloakError as e:
return False, f"Keycloak API error: {str(e)}"
except Exception as e:
return False, f"Unexpected error: {str(e)}"
def get_requestable_groups(username: str) -> tuple[bool, str]:
try:
kc = get_keycloak_client()
users = kc.get_users(query={"username": username, "exact": True})
if not users:
user_group_names = set()
else:
user_id = users[0]["id"]
user_groups = kc.get_user_groups(user_id=user_id)
user_group_names = {g["name"] for g in user_groups}
groups = kc.get_groups()
valid_groups = [g["name"] for g in groups if g["name"].lower() != "admin"]
if not valid_groups:
return True, "No groups are currently available to request."
group_lines = []
for name in valid_groups:
if name in user_group_names:
group_lines.append(f"* `{name}` *(Already a member)*")
else:
group_lines.append(f"* `{name}`")
group_list = "\n".join(group_lines)
return True, f"**Available Groups:**\n\n{group_list}"
except KeycloakError as e:
return False, f"Keycloak API error: {str(e)}"
except Exception as e:
return False, f"⚠️ Unexpected error: {str(e)}"
# ------------------------------------------------------------------------------
# Matrix Helper Functions & Event Handlers
# ------------------------------------------------------------------------------
client: Client = None # Global client instance
async def send_markdown_message(
room_id: str,
md_text: str,
msgtype: MessageType = MessageType.TEXT,
edit_event_id: str = None,
):
"""Sends a markdown formatted message, optionally replacing an older message."""
html_body = markdown.markdown(md_text, extensions=["extra", "sane_lists", "nl2br"])
content = TextMessageEventContent(
msgtype=msgtype,
body=md_text,
format=Format.HTML,
formatted_body=html_body,
)
if edit_event_id:
content.set_edit(edit_event_id)
return await client.send_message(room_id, content)
async def on_invite(evt: Event):
"""Auto-joins rooms when invited."""
if evt.state_key == client.mxid and evt.content.membership == Membership.INVITE:
log.info(f"Received invite for room {evt.room_id}. Joining...")
await client.join_room(evt.room_id)
async def on_message(evt: Event):
"""Parses incoming text messages for commands."""
if evt.sender == client.mxid:
return
if not hasattr(evt.content, 'body') or not isinstance(evt.content.body, str):
return
body = evt.content.body.strip()
matrix_user = evt.sender
kc_username = matrix_user.split(":")[0].lstrip("@")
if body.startswith("!groups"):
success, msg = get_requestable_groups(kc_username)
await send_markdown_message(
room_id=evt.room_id,
md_text=msg,
msgtype=MessageType.NOTICE if not success else MessageType.TEXT,
)
return
if not body.startswith("!request"):
return
parts = body.split(maxsplit=1)
if len(parts) < 2:
await send_markdown_message(
room_id=evt.room_id,
md_text="Usage: `!request <group_name>` or `!groups`",
msgtype=MessageType.NOTICE,
)
return
requested_group = parts[1].strip()
if requested_group.lower() == "admin":
await send_markdown_message(
room_id=evt.room_id,
md_text="❌ Requests for the **admin** group are not permitted.",
msgtype=MessageType.NOTICE,
)
return
already_member, membership_msg = check_group_membership(kc_username, requested_group)
if already_member or membership_msg:
await send_markdown_message(
room_id=evt.room_id,
md_text=membership_msg,
msgtype=MessageType.NOTICE,
)
return
admin_msg_body = (
f"📋 **Access Request**\n\n"
f"**User:** `{matrix_user}` (KC: `{kc_username}`)\n"
f"**Requested Group:** `{requested_group}`\n\n"
f"React with 👍 to **Approve** or 👎 to **Deny**."
)
event_id = await send_markdown_message(
room_id=ADMIN_ROOM_ID,
md_text=admin_msg_body,
)
if event_id:
pending_requests[event_id] = {
"matrix_user": matrix_user,
"kc_username": kc_username,
"requested_group": requested_group,
"origin_room_id": evt.room_id,
}
await send_markdown_message(
room_id=evt.room_id,
md_text=f"Request for group `{requested_group}` submitted for admin approval.",
msgtype=MessageType.NOTICE,
)
async def on_reaction(evt: Event):
"""Handles admin approval/denial via emoji reactions."""
if evt.room_id != ADMIN_ROOM_ID or evt.sender not in ADMIN_MATRIX_IDS:
return
target_event_id = evt.content.relates_to.event_id
emoji = evt.content.relates_to.key
if not target_event_id or target_event_id not in pending_requests:
return
request = pending_requests[target_event_id]
if "👍" in emoji:
success, msg = add_user_to_kc_group(request["kc_username"], request["requested_group"])
status_text = (
f"✅ **APPROVED** by `{evt.sender}`\n\n"
f"**User:** `{request['matrix_user']}` → **Group:** `{request['requested_group']}`\n"
f"**Result:** {msg}"
)
user_msg = (
f"🎉 Your request for access to group `{request['requested_group']}` has been approved!"
if success else f"⚠️ Approval failed: {msg}"
)
await send_markdown_message(
room_id=request["origin_room_id"],
md_text=user_msg,
msgtype=MessageType.NOTICE,
)
elif "👎" in emoji:
status_text = (
f"❌ **DENIED** by `{evt.sender}`\n\n"
f"**User:** `{request['matrix_user']}` → **Group:** `{request['requested_group']}`"
)
await send_markdown_message(
room_id=request["origin_room_id"],
md_text=f"Your request for group `{request['requested_group']}` was denied.",
msgtype=MessageType.NOTICE,
)
else:
return
await send_markdown_message(
room_id=ADMIN_ROOM_ID,
md_text=status_text,
edit_event_id=target_event_id,
)
del pending_requests[target_event_id]
# ------------------------------------------------------------------------------
# Main Execution & Setup
# ------------------------------------------------------------------------------
async def main():
global client
# Load session if it exists
session_data = {}
if os.path.exists(SESSION_FILE):
try:
with open(SESSION_FILE, "r") as f:
session_data = json.load(f)
except Exception as e:
log.warning(f"Could not load session file: {e}")
# Initialize client with persisted token and device ID if available
client = Client(
base_url=MATRIX_HOMESERVER,
mxid=MATRIX_BOT_USER,
access_token=session_data.get("access_token"),
device_id=session_data.get("device_id", "KEYCLOAK_BOT_DEVICE")
)
# Initialize stores using Pickle and Custom Memory State Stores
state_store = CustomMemoryStateStore(client.mxid)
crypto_store = PickleCryptoStore(CRYPTO_STORE_FILE)
client.state_store = state_store
client.crypto = OlmMachine(client, crypto_store, state_store)
# Authenticate if access token is missing
if not client.access_token:
log.info("Logging in to Matrix homeserver...")
await client.login(
password=MATRIX_BOT_PASSWORD,
device_name="KEYCLOAK_ACCESS_BOT",
device_id="KEYCLOAK_BOT_DEVICE"
)
# Save session data for future restarts
session_dir = os.path.dirname(SESSION_FILE)
if session_dir:
os.makedirs(session_dir, exist_ok=True)
with open(SESSION_FILE, "w") as f:
json.dump({
"access_token": client.access_token,
"device_id": client.device_id
}, f)
# Start the crypto machine
await client.crypto.start()
# Bootstrap Cross-Signing and Key Recovery using the environment variable
recovery_key = os.getenv("MATRIX_RECOVERY_KEY")
if recovery_key and client.crypto:
try:
log.info("Verifying recovery key and bootstrapping cross-signing...")
await client.crypto.verify_recovery_key(recovery_key)
if not await client.crypto.has_cross_signing_keys():
await client.crypto.bootstrap_cross_signing(
auth_data={"password": MATRIX_BOT_PASSWORD}
)
log.info("Cross-signing and key recovery initialized successfully.")
except Exception as e:
log.error(f"Failed to initialize cross-signing/recovery: {e}")
# Register Event Handlers
client.add_event_handler(EventType.ROOM_MEMBER, on_invite)
client.add_event_handler(EventType.ROOM_MESSAGE, on_message)
client.add_event_handler(EventType.REACTION, on_reaction)
# Start Sync Loop
log.info("Starting Mautrix sync loop...")
await client.start(None)
if __name__ == "__main__":
try:
asyncio.run(main())
except KeyboardInterrupt:
print("\nBot stopped by user.")
sys.exit(0)