refactor: extensibility overhaul — handler registry, permission checker, sidebar decomposition, DX improvements

Server:
- Unified permission checker (permissions/checker.go) replaces 3 duplicated implementations
- WS handler registry pattern (ws/registry.go) replaces monolithic switch (747→184 lines)
- Split handlers into domain files: handlers_chat.go, handlers_presence.go, handlers_reaction.go
- Shared message type constants (ws/message_types.go) — no more string literals
- Admin API split into helpers.go, types.go, middleware.go (api.go now 61 lines)
- Dev seed script (scripts/seed.go) with -confirm-dev safety flag
- Air hot reload config (.air.toml)
- Fix: DM attachment permission now uses participant check, not role check
- Fix: Typing broadcast now checks ReadMessages permission for non-DM channels

Client:
- Extract preferences to @lib/preferences.ts (fixes lib→component dependency)
- Extract roles to dedicated roles.store.ts (was mixed into channels store)
- Decompose SidebarArea (921→598 lines) into 4 sub-components
- Shared modal factory (lib/modalFactory.ts) with tests
- Global showToast() helper (lib/toast.ts) — 18 call sites migrated
- Protocol type constants (lib/protocolTypes.ts) synced with server
- Remove 38 unnecessary type casts across 17 files
- Component test harness (tests/helpers/test-harness.ts) with 8 tests
- Fix: DM section "View All" respects collapsed state
- Fix: Modal onClose fires on external signal abort
- Fix: savePref wrapped in try/catch for quota exceeded
- Fix: loadPref null guard added

Triple-reviewed: Claude code-review agent + OpenAI Codex CLI + GitHub Copilot
This commit is contained in:
jevb
2026-03-29 12:19:08 +02:00
parent 2d878fd8cb
commit 39658e919b
53 changed files with 3246 additions and 919 deletions
@@ -92,7 +92,7 @@ export function createCreateChannelModal(
type: "text",
placeholder: isVoiceCategory(category) ? "lounge" : "general",
"data-testid": "channel-name-input",
}) as HTMLInputElement;
});
appendChildren(nameGroup, nameLabel, nameInput);
// Channel type
@@ -101,7 +101,7 @@ export function createCreateChannelModal(
const typeSelect = createElement("select", {
class: "form-input",
"data-testid": "channel-type-select",
}) as HTMLSelectElement;
});
for (const t of allowedTypes) {
const opt = createElement(
@@ -68,7 +68,7 @@ export function createEditChannelModal(
type: "text",
value: channelName,
"data-testid": "edit-channel-name-input",
}) as HTMLInputElement;
});
nameInput.value = channelName;
appendChildren(nameGroup, nameLabel, nameInput);
@@ -93,20 +93,20 @@ export function createFileUpload(options: FileUploadOptions): FileUploadComponen
dropzone = createElement("div", { class: "file-upload__dropzone file-upload__dropzone--hidden" });
appendChildren(dropzone, createElement("span", { class: "file-upload__droptext" }, "Drop files here"));
fileInput = createElement("input", { class: "file-upload__input", type: "file" }) as HTMLInputElement;
fileInput = createElement("input", { class: "file-upload__input", type: "file" });
fileInput.style.display = "none";
preview = createElement("div", { class: "file-upload__preview file-upload__preview--hidden" });
thumb = createElement("img", { class: "file-upload__thumb" }) as HTMLImageElement;
thumb = createElement("img", { class: "file-upload__thumb" });
thumb.style.display = "none";
thumb.alt = "";
nameSpan = createElement("span", { class: "file-upload__name" }) as HTMLSpanElement;
sizeSpan = createElement("span", { class: "file-upload__size" }) as HTMLSpanElement;
nameSpan = createElement("span", { class: "file-upload__name" });
sizeSpan = createElement("span", { class: "file-upload__size" });
const progressContainer = createElement("div", { class: "file-upload__progress" });
progressBar = createElement("div", { class: "file-upload__progress-bar" });
progressBar.style.width = "0%";
appendChildren(progressContainer, progressBar);
cancelBtn = createElement("button", { class: "file-upload__cancel", type: "button" }) as HTMLButtonElement;
cancelBtn = createElement("button", { class: "file-upload__cancel", type: "button" });
cancelBtn.appendChild(createIcon("x", 14));
appendChildren(preview, thumb, nameSpan, sizeSpan, progressContainer, cancelBtn);
@@ -198,7 +198,7 @@ export function createMessageInput(
const img = createElement("img", {
class: "attachment-preview-img",
alt: file.name,
}) as HTMLImageElement;
});
item.appendChild(img);
readFileAsDataUrl(file).then((dataUrl) => {
img.src = dataUrl;
@@ -319,7 +319,7 @@ export function createMessageInput(
type: "file",
style: "display: none;",
accept: "image/*,video/*,audio/*,.pdf,.txt,.zip,.rar,.7z",
}) as HTMLInputElement;
});
fileInput.addEventListener("change", () => {
const file = fileInput.files?.[0];
if (file != null) {
@@ -497,7 +497,7 @@ export function createMessageList(options: MessageListOptions): MessageListCompo
bottomSpacer = createElement("div", { class: "virtual-spacer-bottom" });
const scrollAnchor = createElement("div", { class: "scroll-anchor" });
scrollToBottomBtn = createElement("button", { class: "scroll-to-bottom-btn" }) as HTMLButtonElement;
scrollToBottomBtn = createElement("button", { class: "scroll-to-bottom-btn" });
scrollToBottomBtn.textContent = "↓";
scrollToBottomBtn.addEventListener("click", () => {
scrollToBottom();
@@ -155,7 +155,7 @@ export function createQuickSwitcher(options: QuickSwitcherOptions): MountableCom
class: "quick-switcher__input",
type: "text",
placeholder: "Where do you want to go?",
}) as HTMLInputElement;
});
// Results list
resultsDiv = createElement("div", { class: "quick-switcher__results" });
@@ -187,7 +187,7 @@ export function createSearchOverlay(options: SearchOverlayOptions): MountableCom
placeholder: "Search messages...",
"aria-label": "Search messages",
"data-testid": "search-overlay-input",
}) as HTMLInputElement;
});
statusEl = createElement("div", {
class: "search-overlay-status",
@@ -176,7 +176,7 @@ export function createVideoGrid(): VideoGridComponent {
value: "100",
class: "tile-volume-slider",
"aria-label": "Volume",
}) as HTMLInputElement;
});
volumeSlider.addEventListener("input", () => {
currentVolume = Number(volumeSlider.value);
@@ -198,7 +198,7 @@ export function createVideoGrid(): VideoGridComponent {
const muteBtn = createElement("button", {
class: "tile-mute-btn",
"aria-label": "Mute",
}) as HTMLButtonElement;
});
muteBtn.appendChild(volumeIcon());
muteBtn.addEventListener("click", () => {
@@ -236,7 +236,7 @@ export function renderAttachment(att: Attachment): HTMLDivElement {
const img = createElement("img", {
src: cached,
alt: att.filename,
}) as HTMLImageElement;
});
attachLightbox(img);
img.addEventListener("load", () => {
clearReservation();
@@ -253,7 +253,7 @@ export function renderAttachment(att: Attachment): HTMLDivElement {
const img = createElement("img", {
src: dataUrl,
alt: att.filename,
}) as HTMLImageElement;
});
attachLightbox(img);
img.addEventListener("load", () => {
clearReservation();
@@ -259,7 +259,7 @@ export function openImageLightbox(src: string, alt: string): void {
const overlay = createElement("div", { class: "image-lightbox" });
const imgWrap = createElement("div", { class: "image-lightbox-wrap" });
const img = createElement("img", { src, alt }) as HTMLImageElement;
const img = createElement("img", { src, alt });
imgWrap.appendChild(img);
overlay.appendChild(imgWrap);
@@ -173,7 +173,7 @@ function buildCacheRow(
const descEl = createElement("div", { class: "setting-desc" }, desc);
appendChildren(info, labelEl, descEl);
const btn = createElement("button", { class: "ac-btn" }, btnText) as HTMLButtonElement;
const btn = createElement("button", { class: "ac-btn" }, btnText);
btn.addEventListener("click", () => { onClick(btn); }, { signal });
appendChildren(row, info, btn);
@@ -108,7 +108,7 @@ export function buildAppearanceTab(signal: AbortSignal): HTMLDivElement {
placeholder: "5865f2",
value: currentAccent.replace("#", ""),
style: "width:120px",
}) as HTMLInputElement;
});
for (const color of ACCENT_PRESETS) {
const swatch = createElement("div", {
@@ -194,7 +194,7 @@ function buildVoiceAudioTabInner(signal: AbortSignal, registerMic: MicRegistrar,
const qualitySelect = createElement("select", {
class: "form-input",
style: "width:100%;margin-bottom:16px",
}) as HTMLSelectElement;
});
const qualityOptions: Array<[string, string]> = [
["low", "Low (360p cam / 720p screen)"],
["medium", "Medium (720p)"],
@@ -318,7 +318,7 @@ function buildVoiceAudioTabInner(signal: AbortSignal, registerMic: MicRegistrar,
previewVideo.srcObject = stream;
} catch (err) {
const msg = err instanceof Error ? err.message : "Camera unavailable";
previewErrorEl = createElement("div", { class: "setting-desc" }, msg) as HTMLDivElement;
previewErrorEl = createElement("div", { class: "setting-desc" }, msg);
previewWrap.appendChild(previewErrorEl);
}
})();
+27 -26
View File
@@ -52,6 +52,7 @@ import type { DmChannelPayload } from "./types";
import { handleVoiceToken } from "@lib/livekitSession";
import { notifyIncomingMessage } from "./notifications";
import { createLogger } from "./logger";
import { ServerMessageType as S } from "./protocolTypes";
const log = createLogger("dispatcher");
@@ -85,7 +86,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Auth ──────────────────────────────────────────────
unsubs.push(
ws.on("auth_ok", (payload) => {
ws.on(S.AUTH_OK, (payload) => {
setAuth(
authStore.getState().token ?? "",
payload.user,
@@ -96,7 +97,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
);
unsubs.push(
ws.on("auth_error", (payload) => {
ws.on(S.AUTH_ERROR, (payload) => {
log.error("Auth failed", { message: payload.message });
setTransientError(payload.message);
clearAuth();
@@ -106,7 +107,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Ready (initial state dump) ────────────────────────
unsubs.push(
ws.on("ready", (payload) => {
ws.on(S.READY, (payload) => {
setChannels(payload.channels);
setRoles(payload.roles ?? []);
setMembers(payload.members);
@@ -139,14 +140,14 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── DM Channels ─────────────────────────────────────
unsubs.push(
ws.on("dm_channel_open", (payload) => {
ws.on(S.DM_CHANNEL_OPEN, (payload) => {
log.info("DM channel opened", { channelId: payload.channel_id });
addDmChannel(mapDmPayload(payload));
}),
);
unsubs.push(
ws.on("dm_channel_close", (payload) => {
ws.on(S.DM_CHANNEL_CLOSE, (payload) => {
log.info("DM channel closed", { channelId: payload.channel_id });
removeDmChannel(payload.channel_id);
}),
@@ -155,7 +156,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Chat Messages ─────────────────────────────────────
unsubs.push(
ws.on("chat_message", (payload) => {
ws.on(S.CHAT_MESSAGE, (payload) => {
log.debug("chat_message received", {
id: payload.id,
channelId: payload.channel_id,
@@ -208,19 +209,19 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
);
unsubs.push(
ws.on("chat_edited", (payload) => {
ws.on(S.CHAT_EDITED, (payload) => {
editMessage(payload);
}),
);
unsubs.push(
ws.on("chat_deleted", (payload) => {
ws.on(S.CHAT_DELETED, (payload) => {
deleteMessage(payload);
}),
);
unsubs.push(
ws.on("chat_send_ok", (payload, id) => {
ws.on(S.CHAT_SEND_OK, (payload, id) => {
if (id) {
confirmSend(id, payload.message_id, payload.timestamp);
}
@@ -230,7 +231,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Reactions ───────────────────────────────────────────
unsubs.push(
ws.on("reaction_update", (payload) => {
ws.on(S.REACTION_UPDATE, (payload) => {
const userId = authStore.getState().user?.id ?? 0;
updateReaction(payload, userId);
}),
@@ -239,7 +240,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Typing ────────────────────────────────────────────
unsubs.push(
ws.on("typing", (payload) => {
ws.on(S.TYPING, (payload) => {
setTyping(payload.channel_id, payload.user_id);
}),
);
@@ -247,7 +248,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Presence ──────────────────────────────────────────
unsubs.push(
ws.on("presence", (payload) => {
ws.on(S.PRESENCE, (payload) => {
updatePresence(payload.user_id, payload.status);
}),
);
@@ -255,19 +256,19 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Channels ──────────────────────────────────────────
unsubs.push(
ws.on("channel_create", (payload) => {
ws.on(S.CHANNEL_CREATE, (payload) => {
addChannel(payload);
}),
);
unsubs.push(
ws.on("channel_update", (payload) => {
ws.on(S.CHANNEL_UPDATE, (payload) => {
updateChannel(payload);
}),
);
unsubs.push(
ws.on("channel_delete", (payload) => {
ws.on(S.CHANNEL_DELETE, (payload) => {
// If the deleted channel is the active one, redirect to the first text channel.
const activeId = channelsStore.select((s) => s.activeChannelId);
removeChannel(payload.id);
@@ -286,28 +287,28 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Members ───────────────────────────────────────────
unsubs.push(
ws.on("member_join", (payload) => {
ws.on(S.MEMBER_JOIN, (payload) => {
log.info("Member joined", { userId: payload.user.id, username: payload.user.username });
addMember(payload);
}),
);
unsubs.push(
ws.on("member_leave", (payload) => {
ws.on(S.MEMBER_LEAVE, (payload) => {
log.info("Member left", { userId: payload.user_id });
removeMember(payload.user_id);
}),
);
unsubs.push(
ws.on("member_ban", (payload) => {
ws.on(S.MEMBER_BAN, (payload) => {
log.info("Member banned", { userId: payload.user_id });
removeMember(payload.user_id);
}),
);
unsubs.push(
ws.on("member_update", (payload) => {
ws.on(S.MEMBER_UPDATE, (payload) => {
log.info("Member role updated", { userId: payload.user_id, role: payload.role });
updateMemberRole(payload.user_id, payload.role);
}),
@@ -316,7 +317,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Voice ─────────────────────────────────────────────
unsubs.push(
ws.on("voice_state", (payload) => {
ws.on(S.VOICE_STATE, (payload) => {
updateVoiceState(payload);
// Auto-join voice channel if the event is for the current user
const currentUserId = authStore.getState().user?.id ?? 0;
@@ -327,7 +328,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
);
unsubs.push(
ws.on("voice_leave", (payload) => {
ws.on(S.VOICE_LEAVE, (payload) => {
removeVoiceUser(payload);
// Clear local voice state if the current user was removed (kick/disconnect)
const currentUserId = authStore.getState().user?.id ?? 0;
@@ -338,19 +339,19 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
);
unsubs.push(
ws.on("voice_config", (payload) => {
ws.on(S.VOICE_CONFIG, (payload) => {
setVoiceConfig(payload);
}),
);
unsubs.push(
ws.on("voice_speakers", (payload) => {
ws.on(S.VOICE_SPEAKERS, (payload) => {
setSpeakers(payload);
}),
);
unsubs.push(
ws.on("voice_token", (payload) => {
ws.on(S.VOICE_TOKEN, (payload) => {
void handleVoiceToken(payload.token, payload.url, payload.channel_id, payload.direct_url);
}),
);
@@ -358,7 +359,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
// ── Server Events ─────────────────────────────────────
unsubs.push(
ws.on("server_restart", (payload) => {
ws.on(S.SERVER_RESTART, (payload) => {
log.warn("Server restarting", {
reason: payload.reason,
delaySeconds: payload.delay_seconds,
@@ -368,7 +369,7 @@ export function wireDispatcher(ws: WsClient): DispatcherCleanup {
);
unsubs.push(
ws.on("error", (payload) => {
ws.on(S.ERROR, (payload) => {
log.error("Server error", {
code: payload.code,
message: payload.message,
+138
View File
@@ -0,0 +1,138 @@
/**
* Shared modal overlay factory.
* Creates a modal with backdrop, optional click-outside and Escape key
* dismissal, and clean lifecycle management via AbortController.
*
* CSS classes match the existing project convention:
* - div.modal-overlay.visible (backdrop)
* - div.modal (content container)
*/
import { createElement } from "./dom";
export interface ModalOptions {
/** The content element to place inside the modal container. */
readonly content: HTMLElement;
/** Called when the modal is closed (backdrop click, Escape, or programmatic). */
readonly onClose?: () => void;
/** Close when the backdrop is clicked. Default: true. */
readonly closeOnBackdrop?: boolean;
/** Close when the Escape key is pressed. Default: true. */
readonly closeOnEscape?: boolean;
/** Additional CSS class on the .modal container (e.g. "dm-member-picker-modal"). */
readonly className?: string;
/** Additional attributes on the overlay element (e.g. data-testid). */
readonly overlayAttrs?: Readonly<Record<string, string>>;
/** AbortSignal for automatic cleanup when the parent component is destroyed. */
readonly signal?: AbortSignal;
}
export interface ModalInstance {
/** The overlay element (outermost). */
readonly overlay: HTMLElement;
/** The modal container element (inner). */
readonly modal: HTMLElement;
/** Hide the modal (removes visible class). */
close(): void;
/** Remove the modal from the DOM and clean up all listeners. */
destroy(): void;
}
/**
* Create and append a modal overlay to the given container (default: document.body).
* Returns a ModalInstance for lifecycle control.
*/
export function createModal(
options: ModalOptions,
container: Element = document.body,
): ModalInstance {
const {
content,
onClose,
closeOnBackdrop = true,
closeOnEscape = true,
className,
overlayAttrs,
signal,
} = options;
const ac = new AbortController();
// Build overlay
const overlayBaseAttrs: Record<string, string> = {
class: "modal-overlay visible",
};
if (overlayAttrs !== undefined) {
Object.assign(overlayBaseAttrs, overlayAttrs);
}
const overlay = createElement("div", overlayBaseAttrs);
// Build modal container
const modalClass = className !== undefined
? `modal ${className}`
: "modal";
const modal = createElement("div", { class: modalClass });
modal.appendChild(content);
overlay.appendChild(modal);
let closed = false;
function handleClose(): void {
if (closed) return;
closed = true;
overlay.remove();
ac.abort();
if (onClose !== undefined) {
onClose();
}
}
// Backdrop click
if (closeOnBackdrop) {
overlay.addEventListener(
"click",
(e) => {
if (e.target === overlay) {
handleClose();
}
},
{ signal: ac.signal },
);
}
// Escape key
if (closeOnEscape) {
document.addEventListener(
"keydown",
(e: KeyboardEvent) => {
if (e.key === "Escape") {
handleClose();
}
},
{ signal: ac.signal },
);
}
// If an external signal is provided, clean up when it aborts
if (signal !== undefined) {
signal.addEventListener("abort", () => {
if (!closed) {
closed = true;
overlay.remove();
onClose?.();
if (!ac.signal.aborted) {
ac.abort();
}
}
}, { signal: ac.signal });
}
container.appendChild(overlay);
return {
overlay,
modal,
close: handleClose,
destroy: handleClose,
};
}
@@ -0,0 +1,41 @@
/**
* Preference persistence helpers.
*
* Moved here from `@components/settings/helpers` so that `lib/` modules can
* depend on these utilities without importing from the component layer.
*/
// ---------------------------------------------------------------------------
// Constants
// ---------------------------------------------------------------------------
export const STORAGE_PREFIX = "owncord:settings:";
// ---------------------------------------------------------------------------
// Preference helpers
// ---------------------------------------------------------------------------
export function loadPref<T>(key: string, fallback: T): T {
try {
const raw = localStorage.getItem(STORAGE_PREFIX + key);
if (raw === null) return fallback;
const parsed: unknown = JSON.parse(raw);
// Basic typeof guard against corrupted localStorage (covers boolean,
// number, string fallbacks used by current call sites).
if (parsed === null || typeof parsed !== typeof fallback) return fallback;
return parsed as T;
} catch {
return fallback;
}
}
export function savePref(key: string, value: unknown): void {
try {
localStorage.setItem(STORAGE_PREFIX + key, JSON.stringify(value));
// Dispatch a custom event so same-window listeners can invalidate caches.
// The native `storage` event only fires for cross-tab changes.
window.dispatchEvent(new CustomEvent("owncord:pref-change", { detail: { key } }));
} catch {
// localStorage may throw on quota exceeded or when storage is disabled.
}
}
@@ -0,0 +1,84 @@
// Shared WebSocket protocol message type constants.
// Generated from docs/protocol-schema.json — single source of truth for
// both Server (Go) and Client (TypeScript).
//
// Usage: import { MessageType } from "@lib/protocolTypes";
// ws.send({ type: MessageType.CHAT_SEND, payload: { ... } });
// ---------------------------------------------------------------------------
// Server → Client message types
// ---------------------------------------------------------------------------
export const ServerMessageType = {
AUTH_OK: "auth_ok",
AUTH_ERROR: "auth_error",
READY: "ready",
CHAT_MESSAGE: "chat_message",
CHAT_SEND_OK: "chat_send_ok",
CHAT_EDITED: "chat_edited",
CHAT_DELETED: "chat_deleted",
REACTION_UPDATE: "reaction_update",
TYPING: "typing",
PRESENCE: "presence",
CHANNEL_CREATE: "channel_create",
CHANNEL_UPDATE: "channel_update",
CHANNEL_DELETE: "channel_delete",
VOICE_STATE: "voice_state",
VOICE_LEAVE: "voice_leave",
VOICE_CONFIG: "voice_config",
VOICE_TOKEN: "voice_token",
VOICE_SPEAKERS: "voice_speakers",
MEMBER_JOIN: "member_join",
MEMBER_LEAVE: "member_leave",
MEMBER_UPDATE: "member_update",
MEMBER_BAN: "member_ban",
SERVER_RESTART: "server_restart",
ERROR: "error",
// Extensions (not in protocol-schema.json but used in practice)
PONG: "pong",
DM_CHANNEL_OPEN: "dm_channel_open",
DM_CHANNEL_CLOSE: "dm_channel_close",
} as const;
export type ServerMessageTypeValue =
(typeof ServerMessageType)[keyof typeof ServerMessageType];
// ---------------------------------------------------------------------------
// Client → Server message types
// ---------------------------------------------------------------------------
export const ClientMessageType = {
AUTH: "auth",
CHAT_SEND: "chat_send",
CHAT_EDIT: "chat_edit",
CHAT_DELETE: "chat_delete",
REACTION_ADD: "reaction_add",
REACTION_REMOVE: "reaction_remove",
TYPING_START: "typing_start",
CHANNEL_FOCUS: "channel_focus",
PRESENCE_UPDATE: "presence_update",
VOICE_JOIN: "voice_join",
VOICE_LEAVE: "voice_leave",
VOICE_MUTE: "voice_mute",
VOICE_DEAFEN: "voice_deafen",
VOICE_CAMERA: "voice_camera",
VOICE_SCREENSHARE: "voice_screenshare",
PING: "ping",
// Extension (not in protocol-schema.json but used in practice)
VOICE_TOKEN_REFRESH: "voice_token_refresh",
} as const;
export type ClientMessageTypeValue =
(typeof ClientMessageType)[keyof typeof ClientMessageType];
// ---------------------------------------------------------------------------
// Unified MessageType — all message types in one object for convenience
// ---------------------------------------------------------------------------
export const MessageType = {
...ServerMessageType,
...ClientMessageType,
} as const;
export type MessageTypeValue =
(typeof MessageType)[keyof typeof MessageType];
+37
View File
@@ -0,0 +1,37 @@
/**
* Global toast helper — eliminates verbose `toast?.show()` plumbing.
*
* Call `initToast(container)` once at app startup (MainPage mount).
* Then import `showToast` anywhere to display notifications.
*/
import type { ToastContainer, ToastType } from "@components/Toast";
let instance: ToastContainer | null = null;
/**
* Register the app-wide ToastContainer. Called once during MainPage mount.
* Subsequent calls replace the previous instance (for hot-reload safety).
*/
export function initToast(container: ToastContainer): void {
instance = container;
}
/**
* Clear the registered instance (called on MainPage destroy).
*/
export function teardownToast(): void {
instance = null;
}
/**
* Show a toast notification globally. No-ops silently if the toast
* container has not been initialized yet.
*/
export function showToast(
message: string,
type: ToastType = "info",
durationMs?: number,
): void {
instance?.show(message, type, durationMs);
}
+12 -9
View File
@@ -14,6 +14,7 @@ import type { ServerBannerControl } from "@components/ServerBanner";
import { createSettingsOverlay } from "@components/SettingsOverlay";
import { createToastContainer } from "@components/Toast";
import type { ToastContainer } from "@components/Toast";
import { initToast, teardownToast, showToast } from "@lib/toast";
import { authStore, clearAuth, updateUser } from "@stores/auth.store";
import { closeSettings } from "@stores/ui.store";
import { updatePresence } from "@stores/members.store";
@@ -213,10 +214,10 @@ export function createMainPage(options: MainPageOptions): MountableComponent {
onChangePassword: async (oldPassword, newPassword) => {
try {
await api.changePassword(oldPassword, newPassword);
toast?.show("Password changed successfully", "success");
showToast("Password changed successfully", "success");
} catch (err) {
const msg = err instanceof Error ? err.message : "Failed to change password";
toast?.show(msg, "error");
showToast(msg, "error");
throw err;
}
},
@@ -224,10 +225,10 @@ export function createMainPage(options: MainPageOptions): MountableComponent {
try {
const updated = await api.updateProfile({ username });
updateUser({ username: updated.username });
toast?.show("Profile updated", "success");
showToast("Profile updated", "success");
} catch (err) {
const msg = err instanceof Error ? err.message : "Failed to update profile";
toast?.show(msg, "error");
showToast(msg, "error");
throw err;
}
},
@@ -235,7 +236,7 @@ export function createMainPage(options: MainPageOptions): MountableComponent {
onDeleteAccount: async (password) => {
await api.deleteAccount(password);
clearAuth();
toast?.show("Account deleted successfully", "success");
showToast("Account deleted successfully", "success");
},
onStatusChange: (status) => {
const userId = getCurrentUserId();
@@ -256,11 +257,12 @@ export function createMainPage(options: MainPageOptions): MountableComponent {
toast = createToastContainer();
toast.mount(root);
children.push(toast);
initToast(toast);
// Message loading controller
msgCtrl = createMessageController({
api,
showError: (msg) => toast?.show(msg, "error"),
showError: (msg) => showToast(msg, "error"),
});
// Reaction controller
@@ -268,7 +270,7 @@ export function createMainPage(options: MainPageOptions): MountableComponent {
ws,
reactionsLimiter: limiters.reactions,
getChannelId: () => channelCtrl?.currentChannelId ?? 0,
showError: (msg) => toast?.show(msg, "error"),
showError: (msg) => showToast(msg, "error"),
});
// Channel controller (mount/destroy MessageList, TypingIndicator, MessageInput per channel)
@@ -279,7 +281,7 @@ export function createMainPage(options: MainPageOptions): MountableComponent {
pendingDeleteManager,
reactionCtrl: reactionCtrl!,
typingLimiter: limiters.typing,
showToast: (msg, type) => toast?.show(msg, type as "success" | "error" | "info"),
showToast: (msg, type) => showToast(msg, type as "success" | "error" | "info"),
getCurrentUserId,
slots: {
messagesSlot: chatAreaResult.slots.messagesSlot,
@@ -291,7 +293,7 @@ export function createMainPage(options: MainPageOptions): MountableComponent {
});
// Wire voice error callback to toast
setVoiceOnError((msg) => toast?.show(msg, "error"));
setVoiceOnError((msg) => showToast(msg, "error"));
// Wire remote video callbacks to video grid
const SCREENSHARE_TILE_ID_OFFSET = 1_000_000;
@@ -383,6 +385,7 @@ export function createMainPage(options: MainPageOptions): MountableComponent {
function destroy(): void {
log.info("MainPage destroying");
try {
teardownToast();
// Full voice cleanup — tears down room, callbacks, ws ref, serverHost.
// Prevents stale module-level state persisting across logout/reconnect cycles.
voiceCleanupAll();
@@ -180,15 +180,15 @@ export function createLoginForm(opts: LoginFormOptions): LoginFormApi {
// Host
const hostGroup = buildFormGroup("host", "Server Address", "text", "localhost:8443");
hostInput = qs("input", hostGroup) as HTMLInputElement;
hostInput = qs("input", hostGroup)!;
// Username
const usernameGroup = buildFormGroup("username", "Username", "text", "");
usernameInput = qs("input", usernameGroup) as HTMLInputElement;
usernameInput = qs("input", usernameGroup)!;
// Password
const passwordGroup = buildFormGroup("password", "Password", "password", "");
passwordInput = qs("input", passwordGroup) as HTMLInputElement;
passwordInput = qs("input", passwordGroup)!;
// Remember password checkbox
const rememberGroup = createElement("div", { class: "form-group remember-password-group" });
@@ -205,7 +205,7 @@ export function createLoginForm(opts: LoginFormOptions): LoginFormApi {
// Invite code (register only, hidden by default)
inviteGroup = buildFormGroup("invite", "Invite Code", "text", "");
inviteGroup.classList.add("form-group--hidden");
inviteInput = qs("input", inviteGroup) as HTMLInputElement;
inviteInput = qs("input", inviteGroup)!;
// Submit button
submitBtn = createElement("button", {
@@ -220,7 +220,7 @@ export function createLoginForm(opts: LoginFormOptions): LoginFormApi {
// Toggle mode link
const formSwitch = createElement("div", { class: "form-switch" });
toggleModeBtn = createElement("a", {}, "Need an account? Register") as HTMLAnchorElement;
toggleModeBtn = createElement("a", {}, "Need an account? Register");
formSwitch.appendChild(toggleModeBtn);
appendChildren(form, hostGroup, usernameGroup, passwordGroup, rememberGroup, inviteGroup, submitBtn, formSwitch);
@@ -292,8 +292,8 @@ export function createServerPanel(
}
function handleSave(): void {
const name = (nameInput as HTMLInputElement).value.trim();
const addr = (hostAddrInput as HTMLInputElement).value.trim();
const name = nameInput.value.trim();
const addr = hostAddrInput.value.trim();
if (!name || !addr) return;
onAddProfile!(name, addr);
closeModal();
@@ -317,7 +317,7 @@ export function createServerPanel(
// Mount onto the panel's closest connect-page root
const root = panelEl.closest(".connect-page") ?? document.body;
root.appendChild(overlay);
(nameInput as HTMLInputElement).focus();
nameInput.focus();
}
// ---------------------------------------------------------------------------
@@ -68,7 +68,6 @@ export function createChatArea(opts: ChatAreaOptions): ChatAreaResult {
const pinnedCtrl = createPinnedPanelController({
api,
getRoot,
getToast,
getCurrentChannelId: () => getChannelCtrl()?.currentChannelId ?? null,
onJumpToMessage: (msgId: number) => {
const ctrl = getChannelCtrl();
@@ -81,7 +80,6 @@ export function createChatArea(opts: ChatAreaOptions): ChatAreaResult {
const searchCtrl = createSearchOverlayController({
api,
getRoot,
getToast,
getCurrentChannelId: () => getChannelCtrl()?.currentChannelId ?? null,
onJumpToMessage: (_channelId: number, msgId: number) => {
const ctrl = getChannelCtrl();
@@ -102,27 +100,27 @@ export function createChatArea(opts: ChatAreaOptions): ChatAreaResult {
const chatArea = createElement("div", {
class: "chat-area",
"data-testid": "chat-area",
}) as HTMLDivElement;
});
chatArea.appendChild(chatHeader.element);
// --- Slots ---
const messagesSlot = createElement("div", {
class: "messages-slot",
"data-testid": "messages-slot",
}) as HTMLDivElement;
});
const typingSlot = createElement("div", {
class: "typing-slot",
"data-testid": "typing-slot",
}) as HTMLDivElement;
});
const inputSlot = createElement("div", {
class: "input-slot",
"data-testid": "input-slot",
}) as HTMLDivElement;
});
const videoGridSlot = createElement("div", {
class: "video-grid-slot",
"data-testid": "video-grid-slot",
style: "display:none;flex:1;min-height:0",
}) as HTMLDivElement;
});
// --- Video grid ---
const videoGrid = createVideoGrid();
@@ -0,0 +1,105 @@
/**
* MemberPickerModal — a simple modal that lists server members for starting
* a new DM conversation. Uses the shared modal factory for overlay behavior.
*/
import { createElement, setText, appendChildren } from "@lib/dom";
import { createModal } from "@lib/modalFactory";
import type { ModalInstance } from "@lib/modalFactory";
import type { MountableComponent } from "@lib/safe-render";
import { membersStore } from "@stores/members.store";
import { authStore } from "@stores/auth.store";
// ---------------------------------------------------------------------------
// Types
// ---------------------------------------------------------------------------
export interface MemberPickerOptions {
/** Called when the user selects a member. Receives the member's user ID. */
readonly onSelect: (userId: number) => void;
/** Called when the modal is dismissed (cancel or overlay click). */
readonly onClose: () => void;
}
// ---------------------------------------------------------------------------
// createMemberPickerModal
// ---------------------------------------------------------------------------
/**
* Create and mount a member picker modal. Returns a MountableComponent for
* lifecycle management by the caller.
*/
export function createMemberPickerModal(opts: MemberPickerOptions): MountableComponent {
let modalInstance: ModalInstance | null = null;
function mount(container: Element): void {
const members = membersStore.getState().members;
const currentUserId = authStore.getState().user?.id ?? 0;
// Build the content that goes inside the modal
const content = createElement("div", { style: "padding:20px;" });
const title = createElement("h3", {}, "New Direct Message");
const subtitle = createElement("p", { style: "color:var(--text-secondary);font-size:0.85rem;margin:0 0 8px;" },
"Select a member to start a conversation");
const listContainer = createElement("div", {
class: "dm-member-picker-list",
style: "max-height:300px;overflow-y:auto;",
});
for (const member of members.values()) {
if (member.id === currentUserId) continue;
const item = createElement("div", {
class: "dm-member-picker-item channel-item",
style: "cursor:pointer;padding:6px 8px;display:flex;align-items:center;gap:8px;",
});
const avatar = createElement("div", {
class: "dm-avatar",
style: "width:28px;height:28px;border-radius:50%;background:#5865F2;display:flex;align-items:center;justify-content:center;font-size:0.75rem;color:white;flex-shrink:0;",
});
setText(avatar, member.username.charAt(0).toUpperCase());
const nameEl = createElement("span", {}, member.username);
const statusEl = createElement("span", {
style: `font-size:0.75rem;margin-left:auto;color:${member.status === "online" ? "var(--green)" : "var(--text-micro)"};`,
}, member.status);
appendChildren(item, avatar, nameEl, statusEl);
item.addEventListener("click", () => {
if (modalInstance !== null) {
modalInstance.close();
}
opts.onSelect(member.id);
});
listContainer.appendChild(item);
}
const cancelBtn = createElement("button", {
class: "btn btn-secondary",
style: "margin-top:12px;width:100%;",
}, "Cancel");
cancelBtn.addEventListener("click", () => {
if (modalInstance !== null) {
modalInstance.close();
}
});
appendChildren(content, title, subtitle, listContainer, cancelBtn);
modalInstance = createModal(
{
content,
onClose: opts.onClose,
className: "dm-member-picker-modal",
},
container,
);
}
function destroy(): void {
if (modalInstance !== null) {
modalInstance.destroy();
modalInstance = null;
}
}
return { mount, destroy };
}
@@ -13,7 +13,7 @@ import type { InviteResponse } from "@lib/types";
import { createPinnedMessages } from "@components/PinnedMessages";
import type { PinnedMessage } from "@components/PinnedMessages";
import { createSearchOverlay } from "@components/SearchOverlay";
import type { ToastContainer } from "@components/Toast";
import { showToast } from "@lib/toast";
import { setActiveChannel } from "@stores/channels.store";
const log = createLogger("overlays");
@@ -135,7 +135,7 @@ export interface InviteManagerController {
export function createInviteManagerController(opts: {
readonly api: ApiClient;
readonly getRoot: () => HTMLDivElement | null;
readonly getToast: () => ToastContainer | null;
}): InviteManagerController {
let instance: MountableComponent | null = null;
@@ -176,7 +176,7 @@ export function createInviteManagerController(opts: {
onClose: close,
onError: (message: string) => {
log.error(message);
opts.getToast()?.show(message, "error");
showToast(message, "error");
},
});
if (root !== null) {
@@ -184,7 +184,7 @@ export function createInviteManagerController(opts: {
}
} catch (err) {
log.error("Failed to open invite manager", { error: String(err) });
opts.getToast()?.show("Failed to load invites", "error");
showToast("Failed to load invites", "error");
}
}
@@ -203,7 +203,7 @@ export interface PinnedPanelController {
export function createPinnedPanelController(opts: {
readonly api: ApiClient;
readonly getRoot: () => HTMLDivElement | null;
readonly getToast: () => ToastContainer | null;
readonly getCurrentChannelId: () => number | null;
readonly onJumpToMessage?: (messageId: number) => boolean;
}): PinnedPanelController {
@@ -236,7 +236,7 @@ export function createPinnedPanelController(opts: {
if (found) {
close();
} else {
opts.getToast()?.show("Message not in loaded window", "info");
showToast("Message not in loaded window", "info");
}
} else {
close();
@@ -247,7 +247,7 @@ export function createPinnedPanelController(opts: {
close();
}).catch((err: unknown) => {
log.error("Failed to unpin message", { msgId, error: String(err) });
opts.getToast()?.show("Failed to unpin message", "error");
showToast("Failed to unpin message", "error");
});
},
onClose: close,
@@ -257,7 +257,7 @@ export function createPinnedPanelController(opts: {
}
} catch (err) {
log.error("Failed to load pinned messages", { error: String(err) });
opts.getToast()?.show("Failed to load pinned messages", "error");
showToast("Failed to load pinned messages", "error");
}
}
@@ -276,7 +276,7 @@ export interface SearchOverlayController {
export function createSearchOverlayController(opts: {
readonly api: ApiClient;
readonly getRoot: () => HTMLDivElement | null;
readonly getToast: () => ToastContainer | null;
readonly getCurrentChannelId: () => number | null;
readonly onJumpToMessage?: (channelId: number, messageId: number) => boolean;
}): SearchOverlayController {
@@ -304,7 +304,7 @@ export function createSearchOverlayController(opts: {
} catch (err) {
if (err instanceof DOMException && err.name === "AbortError") throw err;
log.error("Search failed", { query, error: String(err) });
opts.getToast()?.show("Search failed", "error");
showToast("Search failed", "error");
throw err;
}
},
@@ -315,7 +315,7 @@ export function createSearchOverlayController(opts: {
requestAnimationFrame(() => {
const found = opts.onJumpToMessage!(result.channel_id, result.message_id);
if (!found) {
opts.getToast()?.show("Message not in loaded history", "info");
showToast("Message not in loaded history", "info");
}
});
}
@@ -99,7 +99,7 @@ export function createSidebarArea(opts: SidebarAreaOptions): SidebarAreaResult {
const sidebarWrapper = createElement("div", {
class: "unified-sidebar",
"data-testid": "unified-sidebar",
}) as HTMLDivElement;
});
// ---------------------------------------------------------------------------
// Server header
@@ -121,7 +121,7 @@ export function createSidebarArea(opts: SidebarAreaOptions): SidebarAreaResult {
serverHeader.appendChild(serverInfoCol);
// Invite button in the server header (proper styled button)
const headerInviteCtrl = createInviteManagerController({ api, getRoot, getToast });
const headerInviteCtrl = createInviteManagerController({ api, getRoot });
const headerInviteBtn = createElement("button", {
class: "sidebar-invite-btn",
title: "Invite people",
@@ -0,0 +1,138 @@
/**
* SidebarDmHelpers — DM-related business logic helpers used by both the
* embedded DM section (channels mode) and the full DM sidebar (dms mode).
*/
import type { ApiClient } from "@lib/api";
import type { ToastContainer } from "@components/Toast";
import type { DmConversation } from "@components/DmSidebar";
import { setSidebarMode, setActiveDmUser } from "@stores/ui.store";
import { channelsStore, setActiveChannel } from "@stores/channels.store";
import type { Channel } from "@stores/channels.store";
import { dmStore, clearDmUnread, addDmChannel } from "@stores/dm.store";
import type { DmChannel } from "@stores/dm.store";
import { membersStore } from "@stores/members.store";
// ---------------------------------------------------------------------------
// Types
// ---------------------------------------------------------------------------
export interface DmHelperDeps {
readonly api: ApiClient;
readonly getToast: () => ToastContainer | null;
readonly getChannelBeforeDm: () => number | null;
readonly setChannelBeforeDm: (id: number | null) => void;
}
// ---------------------------------------------------------------------------
// selectDmConversation
// ---------------------------------------------------------------------------
/**
* Switch the UI to a specific DM conversation. Saves the current non-DM
* channel so it can be restored when the user navigates back.
*/
export function selectDmConversation(
dmChannel: DmChannel,
deps: DmHelperDeps,
): void {
// Save current channel so we can restore it when user clicks "Back"
// Only save if the current channel is a real text/voice channel, not another DM
const currentActive = channelsStore.getState().activeChannelId;
if (currentActive !== null) {
const currentCh = channelsStore.getState().channels.get(currentActive);
if (currentCh !== undefined && currentCh.type !== "dm") {
deps.setChannelBeforeDm(currentActive);
}
}
setActiveDmUser(dmChannel.recipient.id);
setSidebarMode("dms");
clearDmUnread(dmChannel.channelId);
// Add the DM channel to channelsStore so ChannelController can load it
addDmToChannelsStore(dmChannel);
setActiveChannel(dmChannel.channelId);
}
// ---------------------------------------------------------------------------
// addDmToChannelsStore
// ---------------------------------------------------------------------------
/** Ensure a DM channel exists in channelsStore so ChannelController can switch to it. */
export function addDmToChannelsStore(dmChannel: DmChannel): void {
const existing = channelsStore.getState().channels.get(dmChannel.channelId);
// If the channel exists but has an empty name (server sends DMs with name=''),
// update it with the recipient's username
if (existing !== undefined && existing.name !== "") return;
const newChannel: Channel = {
id: dmChannel.channelId,
name: dmChannel.recipient.username,
type: "dm",
category: null,
position: 0,
unreadCount: dmChannel.unreadCount,
lastMessageId: dmChannel.lastMessageId,
};
channelsStore.setState((prev) => {
const next = new Map(prev.channels);
next.set(newChannel.id, newChannel);
return { ...prev, channels: next };
});
}
// ---------------------------------------------------------------------------
// handleCreateDm
// ---------------------------------------------------------------------------
/** Create a DM with a user via the API and switch to it. */
export async function handleCreateDm(
recipientId: number,
deps: DmHelperDeps,
): Promise<void> {
try {
const result = await deps.api.createDm(recipientId);
const member = membersStore.getState().members.get(recipientId);
const dmChannel: DmChannel = {
channelId: result.channel_id,
recipient: {
id: result.recipient.id,
username: result.recipient.username,
avatar: result.recipient.avatar,
status: result.recipient.status ?? member?.status ?? "offline",
},
lastMessageId: null,
lastMessage: "",
lastMessageAt: "",
unreadCount: 0,
};
addDmChannel(dmChannel);
selectDmConversation(dmChannel, deps);
} catch (err) {
const msg = err instanceof Error ? err.message : "Failed to create DM";
deps.getToast()?.show(msg, "error");
}
}
// ---------------------------------------------------------------------------
// buildDmConversations — helper for DM sidebar mode
// ---------------------------------------------------------------------------
/** Build a readonly DmConversation array from DM store state. */
export function buildDmConversations(activeDmUserId: number | null): readonly DmConversation[] {
const dmChannels = dmStore.getState().channels;
return dmChannels.map((dm) => ({
userId: dm.recipient.id,
username: dm.recipient.username,
avatar: dm.recipient.avatar || null,
status: (dm.recipient.status as DmConversation["status"]) ?? "offline",
lastMessage: dm.lastMessage || "No messages yet",
timestamp: dm.lastMessageAt,
unread: dm.unreadCount > 0,
active: dm.recipient.id === activeDmUserId,
}));
}
@@ -0,0 +1,151 @@
/**
* SidebarDmSection — the embedded DM preview section that sits above channels
* in "channels" mode. Shows the top 3 DM conversations, an unread badge,
* a "View all messages" button, and collapse toggle.
*/
import { createElement, setText, clearChildren, appendChildren } from "@lib/dom";
import { dmStore } from "@stores/dm.store";
import type { DmChannel } from "@stores/dm.store";
import { setSidebarMode } from "@stores/ui.store";
// ---------------------------------------------------------------------------
// Types
// ---------------------------------------------------------------------------
export interface SidebarDmSectionOptions {
/** Called when the user clicks a DM entry to open that conversation. */
readonly onSelectDm: (dmChannel: DmChannel) => void;
/** Called when the user clicks the "+" button to create a new DM. */
readonly onNewDm: () => void;
}
export interface SidebarDmSectionResult {
/** The root element to insert into the DOM. */
readonly element: HTMLDivElement;
/** Re-render the DM list from current store state. */
readonly update: () => void;
/** Clean up store subscriptions. */
readonly destroy: () => void;
}
// ---------------------------------------------------------------------------
// Factory
// ---------------------------------------------------------------------------
export function createSidebarDmSection(opts: SidebarDmSectionOptions): SidebarDmSectionResult {
const unsubs: Array<() => void> = [];
// --- Root container ---
const dmSection = createElement("div", { class: "sidebar-dm-section" });
// --- Header ---
const dmHeader = createElement("div", { class: "category" });
const dmArrow = createElement("span", { class: "category-arrow" }, "\u25BC");
const dmLabelEl = createElement("span", { class: "category-name" }, "DIRECT MESSAGES");
const dmUnreadBadge = createElement("span", { class: "dm-header-unread-badge" });
const dmAddBtn = createElement("button", { class: "category-add-btn", title: "New DM" }, "+");
dmAddBtn.style.opacity = "1";
appendChildren(dmHeader, dmArrow, dmLabelEl, dmUnreadBadge, dmAddBtn);
dmSection.appendChild(dmHeader);
// --- DM list ---
let dmCollapsed = false;
const dmList = createElement("div", { class: "category-channels sidebar-dm-list" });
// --- "View All" button ---
const viewAllBtn = createElement("button", {
class: "sidebar-dm-view-all",
}, "View all messages");
viewAllBtn.addEventListener("click", () => {
setSidebarMode("dms");
});
// --- Render logic ---
function renderDmListItems(): void {
clearChildren(dmList);
const dmChannels = dmStore.getState().channels;
const displayChannels = dmChannels.slice(0, 3);
for (const dm of displayChannels) {
const dmItem = createElement("div", {
class: "channel-item",
"data-testid": "dm-entry",
});
const statusColor = dm.recipient.status === "online" ? "var(--green)"
: dm.recipient.status === "idle" ? "var(--yellow)"
: dm.recipient.status === "dnd" ? "var(--red)"
: "var(--text-micro)";
const statusDot = createElement("span", {
style: `display:inline-block;width:8px;height:8px;border-radius:50%;background:${statusColor};flex-shrink:0;`,
});
const name = createElement("span", { class: "ch-name" }, dm.recipient.username);
const parts: Element[] = [statusDot, name];
if (dm.unreadCount > 0) {
const badge = createElement("span", {
class: "dm-unread-badge",
style: "margin-left:auto;background:var(--red);color:white;border-radius:10px;padding:1px 6px;font-size:0.7rem;",
}, String(dm.unreadCount));
parts.push(badge);
}
appendChildren(dmItem, ...parts);
dmItem.addEventListener("click", () => {
opts.onSelectDm(dm);
});
dmList.appendChild(dmItem);
}
// Show/hide "View All" button based on DM count (respect collapsed state)
if (dmChannels.length > 3) {
setText(viewAllBtn, `View all messages (${dmChannels.length})`);
viewAllBtn.style.display = dmCollapsed ? "none" : "";
} else {
viewAllBtn.style.display = "none";
}
// Update total unread badge on the DM header
const totalUnread = dmChannels.reduce((sum, c) => sum + c.unreadCount, 0);
if (totalUnread > 0) {
setText(dmUnreadBadge, String(totalUnread));
dmUnreadBadge.style.display = "";
} else {
dmUnreadBadge.style.display = "none";
}
}
renderDmListItems();
dmSection.appendChild(dmList);
dmSection.appendChild(viewAllBtn);
// --- Store subscription ---
const unsubDmSection = dmStore.subscribeSelector(
(s) => s.channels,
() => { renderDmListItems(); },
);
unsubs.push(unsubDmSection);
// --- Collapse toggle ---
dmHeader.addEventListener("click", () => {
dmCollapsed = !dmCollapsed;
dmHeader.classList.toggle("collapsed", dmCollapsed);
dmArrow.textContent = dmCollapsed ? "\u25B6" : "\u25BC";
dmList.style.display = dmCollapsed ? "none" : "";
viewAllBtn.style.display = dmCollapsed ? "none" : (dmStore.getState().channels.length > 3 ? "" : "none");
});
// --- Add DM button ---
dmAddBtn.addEventListener("click", (e) => {
e.stopPropagation();
opts.onNewDm();
});
return {
element: dmSection,
update: renderDmListItems,
destroy: () => {
for (const unsub of unsubs) {
unsub();
}
},
};
}
@@ -0,0 +1,176 @@
/**
* SidebarMemberSection — the collapsible member list panel that sits below
* channels in "channels" mode. Supports drag-to-resize and persists
* collapsed state and height to localStorage.
*/
import { createElement, appendChildren } from "@lib/dom";
import type { MountableComponent } from "@lib/safe-render";
import { createMemberList } from "@components/MemberList";
import { authStore } from "@stores/auth.store";
import { getRoleIdByName } from "@stores/roles.store";
import type { ApiClient } from "@lib/api";
import type { ToastContainer } from "@components/Toast";
// ---------------------------------------------------------------------------
// Constants
// ---------------------------------------------------------------------------
const LS_KEY_HEIGHT = "owncord:member-list-height";
const LS_KEY_COLLAPSED = "owncord:member-list-collapsed";
// ---------------------------------------------------------------------------
// Types
// ---------------------------------------------------------------------------
export interface SidebarMemberSectionOptions {
readonly api: ApiClient;
readonly getToast: () => ToastContainer | null;
}
export interface SidebarMemberSectionResult {
/** The root element to insert into the DOM. */
readonly element: HTMLDivElement;
/** The member list MountableComponent (for external cleanup tracking). */
readonly memberListComponent: MountableComponent;
/** Clean up event listeners and abort controller. */
readonly destroy: () => void;
}
// ---------------------------------------------------------------------------
// Factory
// ---------------------------------------------------------------------------
export function createSidebarMemberSection(opts: SidebarMemberSectionOptions): SidebarMemberSectionResult {
const { api, getToast } = opts;
const unsubs: Array<() => void> = [];
// --- Container ---
const memberListContainer = createElement("div", {
class: "sidebar-members-section",
"data-testid": "sidebar-members",
});
// --- Header ---
const memberHeader = createElement("div", { class: "category sidebar-members-header" });
const memberArrow = createElement("span", { class: "category-arrow" }, "\u25BC");
const memberLabelEl = createElement("span", { class: "category-name" }, "MEMBERS");
appendChildren(memberHeader, memberArrow, memberLabelEl);
memberListContainer.appendChild(memberHeader);
// --- Resize handle ---
const resizeHandle = createElement("div", { class: "sidebar-resize-handle" });
memberListContainer.appendChild(resizeHandle);
// Restore saved height
const savedHeight = localStorage.getItem(LS_KEY_HEIGHT);
if (savedHeight !== null) {
memberListContainer.style.height = `${savedHeight}px`;
}
// --- Drag-to-resize logic ---
const resizeAbort = new AbortController();
let isDragging = false;
let startY = 0;
let startHeight = 0;
resizeHandle.addEventListener("mousedown", (e: MouseEvent) => {
isDragging = true;
startY = e.clientY;
startHeight = memberListContainer.offsetHeight;
e.preventDefault();
}, { signal: resizeAbort.signal });
document.addEventListener("mousemove", (e: MouseEvent) => {
if (!isDragging) return;
const delta = startY - e.clientY;
const maxH = window.innerHeight * 0.65;
const newHeight = Math.max(80, Math.min(startHeight + delta, maxH));
memberListContainer.style.height = `${newHeight}px`;
}, { signal: resizeAbort.signal });
document.addEventListener("mouseup", () => {
if (!isDragging) return;
isDragging = false;
localStorage.setItem(LS_KEY_HEIGHT, String(memberListContainer.offsetHeight));
}, { signal: resizeAbort.signal });
unsubs.push(() => { resizeAbort.abort(); });
// --- Collapse state ---
const savedCollapsed = localStorage.getItem(LS_KEY_COLLAPSED);
let membersCollapsed = savedCollapsed === "true";
const memberContent = createElement("div", { class: "sidebar-members-content" });
function applyMembersCollapsed(): void {
memberHeader.classList.toggle("collapsed", membersCollapsed);
memberArrow.textContent = membersCollapsed ? "\u25B6" : "\u25BC";
memberContent.style.display = membersCollapsed ? "none" : "";
resizeHandle.style.display = membersCollapsed ? "none" : "";
if (membersCollapsed) {
memberListContainer.style.height = "auto";
} else {
const h = localStorage.getItem(LS_KEY_HEIGHT);
if (h !== null) {
memberListContainer.style.height = `${h}px`;
} else {
memberListContainer.style.height = "";
}
}
}
// Apply initial state
applyMembersCollapsed();
memberHeader.addEventListener("click", () => {
membersCollapsed = !membersCollapsed;
localStorage.setItem(LS_KEY_COLLAPSED, String(membersCollapsed));
applyMembersCollapsed();
});
// --- Member list component ---
const memberList = createMemberList({
currentUserRole: authStore.getState().user?.role ?? "member",
onKick: async (userId, username) => {
try {
await api.adminKickMember(userId);
getToast()?.show(`Kicked ${username}`, "success");
} catch (err) {
const msg = err instanceof Error ? err.message : "Failed to kick member";
getToast()?.show(msg, "error");
}
},
onBan: async (userId, username) => {
try {
await api.adminBanMember(userId);
getToast()?.show(`Banned ${username}`, "success");
} catch (err) {
const msg = err instanceof Error ? err.message : "Failed to ban member";
getToast()?.show(msg, "error");
}
},
onChangeRole: async (userId, username, newRole) => {
const roleId = getRoleIdByName(newRole);
if (roleId === undefined) return;
try {
await api.adminChangeRole(userId, roleId);
getToast()?.show(`Changed ${username}'s role to ${newRole}`, "success");
} catch (err) {
const msg = err instanceof Error ? err.message : "Failed to change role";
getToast()?.show(msg, "error");
}
},
});
memberList.mount(memberContent);
memberListContainer.appendChild(memberContent);
return {
element: memberListContainer,
memberListComponent: memberList,
destroy: () => {
for (const unsub of unsubs) {
unsub();
}
},
};
}
@@ -0,0 +1,29 @@
/**
* Roles store — holds server-wide role definitions.
* Immutable state updates only.
*/
import { createStore } from "@lib/store";
import type { ReadyRole } from "@lib/types";
export interface RolesState {
readonly roles: readonly ReadyRole[];
}
const INITIAL_STATE: RolesState = {
roles: [],
};
export const rolesStore = createStore<RolesState>(INITIAL_STATE);
/** Bulk set roles from the ready payload. */
export function setRoles(roles: readonly ReadyRole[]): void {
rolesStore.setState(() => ({ roles }));
}
/** Look up a role ID by name (case-insensitive). Returns undefined if not found. */
export function getRoleIdByName(name: string): number | undefined {
const roles = rolesStore.getState().roles;
const match = roles.find((r) => r.name.toLowerCase() === name.toLowerCase());
return match?.id;
}
@@ -0,0 +1,92 @@
/**
* Reusable DOM test harness for OwnCord component unit tests.
*
* Eliminates repeated container creation / teardown boilerplate.
*
* @example
* ```ts
* let harness: TestHarness;
*
* beforeEach(() => { harness = createTestHarness(); });
* afterEach(() => { harness.cleanup(); });
*
* it("renders", () => {
* harness.mount(createMyComponent());
* expect(harness.query(".my-class")).not.toBeNull();
* });
* ```
*/
/** Minimal component interface that the harness can mount. */
export interface Mountable {
mount(el: HTMLElement): void;
destroy?(): void;
}
export interface TestHarness {
/** The container div appended to document.body. */
readonly container: HTMLDivElement;
/** Calls `component.mount(container)`. */
mount(component: Mountable): void;
/** Shorthand for `container.querySelector`. */
query<E extends Element = Element>(selector: string): E | null;
/** Shorthand for `container.querySelectorAll`. */
queryAll<E extends Element = Element>(selector: string): NodeListOf<E>;
/** Finds an element by selector and dispatches a click event. Throws if not found. */
click(selector: string): void;
/** Removes the container from the DOM. Safe to call multiple times. */
cleanup(): void;
}
/**
* Creates a fresh DOM container attached to `document.body` and returns
* helper methods for mounting components, querying, and clicking.
*
* Call `cleanup()` in `afterEach` to remove the container.
*/
export function createTestHarness(): TestHarness {
const container = document.createElement("div");
document.body.appendChild(container);
let cleaned = false;
return {
get container(): HTMLDivElement {
return container;
},
mount(component: Mountable): void {
component.mount(container);
},
query<E extends Element = Element>(selector: string): E | null {
return container.querySelector<E>(selector);
},
queryAll<E extends Element = Element>(selector: string): NodeListOf<E> {
return container.querySelectorAll<E>(selector);
},
click(selector: string): void {
const el = container.querySelector(selector) as HTMLElement | null;
if (el === null) {
throw new Error(
`click("${selector}"): no element found in container`,
);
}
el.click();
},
cleanup(): void {
if (!cleaned) {
container.remove();
cleaned = true;
}
},
};
}
@@ -0,0 +1,157 @@
import { describe, it, expect, vi, beforeEach, afterEach } from "vitest";
import { createModal } from "../../src/lib/modalFactory";
describe("createModal", () => {
let container: HTMLDivElement;
beforeEach(() => {
container = document.createElement("div");
document.body.appendChild(container);
});
afterEach(() => {
container.remove();
// Clean up any stray overlays
document.querySelectorAll(".modal-overlay").forEach((el) => el.remove());
});
it("renders with correct structure (overlay > modal > content)", () => {
const content = document.createElement("div");
content.textContent = "Hello";
const inst = createModal({ content }, container);
expect(inst.overlay.classList.contains("modal-overlay")).toBe(true);
expect(inst.overlay.classList.contains("visible")).toBe(true);
expect(inst.modal.classList.contains("modal")).toBe(true);
expect(inst.modal.textContent).toBe("Hello");
expect(container.contains(inst.overlay)).toBe(true);
});
it("applies additional className to modal container", () => {
const content = document.createElement("div");
const inst = createModal({ content, className: "dm-picker" }, container);
expect(inst.modal.classList.contains("modal")).toBe(true);
expect(inst.modal.classList.contains("dm-picker")).toBe(true);
});
it("applies overlay attributes", () => {
const content = document.createElement("div");
const inst = createModal(
{ content, overlayAttrs: { "data-testid": "my-modal" } },
container,
);
expect(inst.overlay.getAttribute("data-testid")).toBe("my-modal");
});
it("backdrop click closes and calls onClose", () => {
const onClose = vi.fn();
const content = document.createElement("div");
const inst = createModal({ content, onClose }, container);
// Click on the overlay itself (not the modal content)
inst.overlay.dispatchEvent(new MouseEvent("click", { bubbles: true }));
expect(onClose).toHaveBeenCalledTimes(1);
expect(container.contains(inst.overlay)).toBe(false);
});
it("clicking inside modal does not close", () => {
const onClose = vi.fn();
const content = document.createElement("div");
content.textContent = "inner";
const inst = createModal({ content, onClose }, container);
// Click on the modal content, not the overlay
inst.modal.dispatchEvent(new MouseEvent("click", { bubbles: true }));
expect(onClose).not.toHaveBeenCalled();
expect(container.contains(inst.overlay)).toBe(true);
});
it("Escape key closes and calls onClose", () => {
const onClose = vi.fn();
const content = document.createElement("div");
createModal({ content, onClose }, container);
document.dispatchEvent(new KeyboardEvent("keydown", { key: "Escape" }));
expect(onClose).toHaveBeenCalledTimes(1);
});
it("close() removes from DOM", () => {
const content = document.createElement("div");
const inst = createModal({ content }, container);
expect(container.contains(inst.overlay)).toBe(true);
inst.close();
expect(container.contains(inst.overlay)).toBe(false);
});
it("destroy() removes from DOM (alias for close)", () => {
const content = document.createElement("div");
const inst = createModal({ content }, container);
inst.destroy();
expect(container.contains(inst.overlay)).toBe(false);
});
it("closeOnBackdrop=false prevents backdrop close", () => {
const onClose = vi.fn();
const content = document.createElement("div");
const inst = createModal(
{ content, onClose, closeOnBackdrop: false },
container,
);
inst.overlay.dispatchEvent(new MouseEvent("click", { bubbles: true }));
expect(onClose).not.toHaveBeenCalled();
expect(container.contains(inst.overlay)).toBe(true);
});
it("closeOnEscape=false prevents Escape close", () => {
const onClose = vi.fn();
const content = document.createElement("div");
createModal(
{ content, onClose, closeOnEscape: false },
container,
);
document.dispatchEvent(new KeyboardEvent("keydown", { key: "Escape" }));
expect(onClose).not.toHaveBeenCalled();
});
it("cleans up when external signal is aborted", () => {
const externalAc = new AbortController();
const content = document.createElement("div");
const inst = createModal(
{ content, signal: externalAc.signal },
container,
);
expect(container.contains(inst.overlay)).toBe(true);
externalAc.abort();
expect(container.contains(inst.overlay)).toBe(false);
});
it("onClose is called only once even with multiple close triggers", () => {
const onClose = vi.fn();
const content = document.createElement("div");
const inst = createModal({ content, onClose }, container);
inst.close();
inst.close();
inst.destroy();
expect(onClose).toHaveBeenCalledTimes(1);
});
});
@@ -11,12 +11,14 @@ const {
mockInviteManagerDestroy,
mockPinnedMessagesMount,
mockPinnedMessagesDestroy,
mockShowToast,
} = vi.hoisted(() => ({
mockLogError: vi.fn(),
mockInviteManagerMount: vi.fn(),
mockInviteManagerDestroy: vi.fn(),
mockPinnedMessagesMount: vi.fn(),
mockPinnedMessagesDestroy: vi.fn(),
mockShowToast: vi.fn(),
}));
vi.mock("@lib/logger", () => ({
@@ -53,6 +55,12 @@ vi.mock("@stores/channels.store", () => ({
setActiveChannel: vi.fn(),
}));
vi.mock("@lib/toast", () => ({
initToast: vi.fn(),
teardownToast: vi.fn(),
showToast: mockShowToast,
}));
// ---------------------------------------------------------------------------
// Imports (after mocks)
// ---------------------------------------------------------------------------
@@ -123,7 +131,7 @@ describe("createInviteManagerController", () => {
const controller = createInviteManagerController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
});
await controller.open();
@@ -141,7 +149,7 @@ describe("createInviteManagerController", () => {
const controller = createInviteManagerController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
});
await controller.open();
@@ -165,7 +173,7 @@ describe("createInviteManagerController", () => {
const controller = createInviteManagerController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
});
await controller.open();
@@ -187,12 +195,12 @@ describe("createInviteManagerController", () => {
const controller = createInviteManagerController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
});
await controller.open();
expect(toast.show).toHaveBeenCalledWith("Failed to load invites", "error");
expect(mockShowToast).toHaveBeenCalledWith("Failed to load invites", "error");
});
});
@@ -216,7 +224,7 @@ describe("createPinnedPanelController", () => {
const controller = createPinnedPanelController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
getCurrentChannelId: () => 42,
});
@@ -235,7 +243,7 @@ describe("createPinnedPanelController", () => {
const controller = createPinnedPanelController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
getCurrentChannelId: () => 42,
});
@@ -251,7 +259,7 @@ describe("createPinnedPanelController", () => {
// Wait for the async error handling to complete
await vi.waitFor(() => {
expect(toast.show).toHaveBeenCalledWith("Failed to unpin message", "error");
expect(mockShowToast).toHaveBeenCalledWith("Failed to unpin message", "error");
});
// Panel should NOT have been destroyed (still open)
@@ -265,7 +273,7 @@ describe("createPinnedPanelController", () => {
const controller = createPinnedPanelController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
getCurrentChannelId: () => 42,
});
@@ -283,7 +291,7 @@ describe("createPinnedPanelController", () => {
});
// No error toast should be shown
expect(toast.show).not.toHaveBeenCalled();
expect(mockShowToast).not.toHaveBeenCalled();
});
it("onJumpToMessage calls provided scroll callback and closes panel", async () => {
@@ -294,7 +302,7 @@ describe("createPinnedPanelController", () => {
const controller = createPinnedPanelController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
getCurrentChannelId: () => 42,
onJumpToMessage: mockScrollToMessage,
});
@@ -319,7 +327,7 @@ describe("createPinnedPanelController", () => {
const controller = createPinnedPanelController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
getCurrentChannelId: () => 42,
onJumpToMessage: mockScrollToMessage,
});
@@ -333,7 +341,7 @@ describe("createPinnedPanelController", () => {
opts.onJumpToMessage(999);
expect(mockScrollToMessage).toHaveBeenCalledWith(999);
expect(toast.show).toHaveBeenCalledWith(
expect(mockShowToast).toHaveBeenCalledWith(
expect.stringContaining("not in"),
"info",
);
@@ -350,12 +358,12 @@ describe("createPinnedPanelController", () => {
const controller = createPinnedPanelController({
api: api as never,
getRoot: () => root,
getToast: () => toast as never,
getCurrentChannelId: () => 42,
});
await controller.toggle();
expect(toast.show).toHaveBeenCalledWith("Failed to load pinned messages", "error");
expect(mockShowToast).toHaveBeenCalledWith("Failed to load pinned messages", "error");
});
});
@@ -0,0 +1,86 @@
import { describe, it, expect, afterEach } from "vitest";
import { createTestHarness, type TestHarness, type Mountable } from "../helpers/test-harness";
describe("createTestHarness", () => {
let harness: TestHarness;
afterEach(() => {
harness?.cleanup();
});
it("appends container to document.body on creation", () => {
harness = createTestHarness();
expect(document.body.contains(harness.container)).toBe(true);
});
it("mount() calls the component's mount method with the container", () => {
harness = createTestHarness();
let mountedOn: HTMLElement | null = null;
const fakeComponent: Mountable = {
mount(el: HTMLElement) {
mountedOn = el;
el.innerHTML = '<span class="test-child">hello</span>';
},
};
harness.mount(fakeComponent);
expect(mountedOn).toBe(harness.container);
expect(harness.container.innerHTML).toContain("test-child");
});
it("query() finds elements within the container", () => {
harness = createTestHarness();
harness.container.innerHTML = '<div class="target">found</div>';
const el = harness.query(".target");
expect(el).not.toBeNull();
expect(el!.textContent).toBe("found");
});
it("queryAll() returns all matching elements within the container", () => {
harness = createTestHarness();
harness.container.innerHTML =
'<span class="item">a</span><span class="item">b</span>';
const els = harness.queryAll(".item");
expect(els.length).toBe(2);
});
it("click() dispatches a click on the matched element", () => {
harness = createTestHarness();
let clicked = false;
const btn = document.createElement("button");
btn.className = "click-me";
btn.addEventListener("click", () => { clicked = true; });
harness.container.appendChild(btn);
harness.click(".click-me");
expect(clicked).toBe(true);
});
it("click() throws when no element matches the selector", () => {
harness = createTestHarness();
expect(() => harness.click(".nonexistent")).toThrow(
'click(".nonexistent"): no element found in container',
);
});
it("cleanup() removes the container from document.body", () => {
harness = createTestHarness();
const container = harness.container;
expect(document.body.contains(container)).toBe(true);
harness.cleanup();
expect(document.body.contains(container)).toBe(false);
});
it("cleanup() is safe to call multiple times", () => {
harness = createTestHarness();
harness.cleanup();
// Should not throw
harness.cleanup();
});
});
+46
View File
@@ -0,0 +1,46 @@
# Air hot-reload config for OwnCord chat server.
# Usage: cd Server && air
# Docs: https://github.com/air-verse/air
root = "."
tmp_dir = "tmp"
[build]
bin = "./chatserver.exe"
cmd = "go build -o chatserver.exe -ldflags \"-s -w\" ."
delay = 1000
exclude_dir = ["tmp", "scripts", "migrations", "data", "admin/static"]
exclude_file = []
exclude_regex = ["_test\\.go$"]
exclude_unchanged = false
follow_symlink = false
full_bin = ""
include_dir = []
include_ext = ["go", "tmpl", "html"]
include_file = []
kill_delay = "3s"
log = "build-errors.log"
poll = false
poll_interval = 0
rerun = false
rerun_delay = 500
send_interrupt = true
stop_on_error = false
[color]
app = ""
build = "yellow"
main = "magenta"
runner = "green"
watcher = "cyan"
[log]
main_only = false
time = false
[misc]
clean_on_exit = true
[screen]
clear_on_rebuild = true
keep_scroll = true
-229
View File
@@ -1,114 +1,13 @@
package admin
import (
"context"
"encoding/json"
"net/http"
"strconv"
"github.com/go-chi/chi/v5"
"github.com/owncord/server/auth"
"github.com/owncord/server/db"
"github.com/owncord/server/permissions"
"github.com/owncord/server/updater"
)
// ─── Context keys ─────────────────────────────────────────────────────────────
// adminContextKey is an unexported type for context keys in the admin package.
type adminContextKey int
const (
// adminUserKey is the context key for the authenticated *db.User.
adminUserKey adminContextKey = iota
// adminSessionKey is the context key for the authenticated *db.Session.
adminSessionKey
)
// ─── Allowed settings keys ────────────────────────────────────────────────────
// allowedSettingKeys is the whitelist of keys that may be written via
// PATCH /admin/api/settings. Derived from the settings table in SCHEMA.md.
var allowedSettingKeys = map[string]struct{}{
"server_name": {},
"server_icon": {},
"motd": {},
"max_upload_bytes": {},
"voice_quality": {},
"require_2fa": {},
"registration_open": {},
"backup_schedule": {},
"backup_retention": {},
}
// HubBroadcaster is the subset of ws.Hub needed by the admin package.
type HubBroadcaster interface {
BroadcastServerRestart(reason string, delaySeconds int)
BroadcastChannelCreate(ch *db.Channel)
BroadcastChannelUpdate(ch *db.Channel)
BroadcastChannelDelete(channelID int64)
BroadcastMemberBan(userID int64)
BroadcastMemberUpdate(userID int64, roleName string)
ClientCount() int
}
// ─── adminUserResponse ──────────────────────────────────────────────────────
// adminUserResponse is the safe public shape returned by user-listing and
// user-patch endpoints. It deliberately excludes PasswordHash and TOTPSecret.
type adminUserResponse struct {
ID int64 `json:"id"`
Username string `json:"username"`
Avatar *string `json:"avatar,omitempty"`
RoleID int64 `json:"role_id"`
RoleName string `json:"role_name"`
Status string `json:"status"`
CreatedAt string `json:"created_at"`
LastSeen *string `json:"last_seen,omitempty"`
Banned bool `json:"banned"`
BanReason *string `json:"ban_reason,omitempty"`
BanExpires *string `json:"ban_expires,omitempty"`
}
// toAdminUserResponse converts a db.UserWithRole to the safe response shape.
func toAdminUserResponse(u db.UserWithRole) adminUserResponse {
return adminUserResponse{
ID: u.ID,
Username: u.Username,
Avatar: u.Avatar,
RoleID: u.RoleID,
RoleName: u.RoleName,
Status: u.Status,
CreatedAt: u.CreatedAt,
LastSeen: u.LastSeen,
Banned: u.Banned,
BanReason: u.BanReason,
BanExpires: u.BanExpires,
}
}
// toAdminUserResponseFromUser converts a plain db.User to the safe response
// shape, resolving the role name via the database.
func toAdminUserResponseFromUser(database *db.DB, u *db.User) adminUserResponse {
roleName := ""
if role, err := database.GetRoleByID(u.RoleID); err == nil && role != nil {
roleName = role.Name
}
return adminUserResponse{
ID: u.ID,
Username: u.Username,
Avatar: u.Avatar,
RoleID: u.RoleID,
RoleName: roleName,
Status: u.Status,
CreatedAt: u.CreatedAt,
LastSeen: u.LastSeen,
Banned: u.Banned,
BanReason: u.BanReason,
BanExpires: u.BanExpires,
}
}
// ─── NewAdminAPI ──────────────────────────────────────────────────────────────
// NewAdminAPI returns a chi router with all /admin/api/* routes. All routes
@@ -160,131 +59,3 @@ func NewAdminAPI(database *db.DB, version string, hub HubBroadcaster, u *updater
return r
}
// ─── Middleware ───────────────────────────────────────────────────────────────
// adminAuthMiddleware validates the Bearer token and requires ADMINISTRATOR.
// On success it stores the *db.User and *db.Session in the request context so
// downstream handlers can retrieve them without re-querying the database.
func adminAuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
token, ok := auth.ExtractBearerToken(r)
if !ok {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "missing or invalid authorization header")
return
}
hash := auth.HashToken(token)
sess, err := database.GetSessionByTokenHash(hash)
if err != nil || sess == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "invalid or expired session")
return
}
if auth.IsSessionExpired(sess.ExpiresAt) {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "session has expired")
return
}
user, err := database.GetUserByID(sess.UserID)
if err != nil || user == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "user not found")
return
}
role, err := database.GetRoleByID(user.RoleID)
if err != nil || role == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "role not found")
return
}
if !permissions.HasAdmin(role.Permissions) {
writeErr(w, http.StatusForbidden, "FORBIDDEN", "administrator permission required")
return
}
ctx := context.WithValue(r.Context(), adminUserKey, user)
ctx = context.WithValue(ctx, adminSessionKey, sess)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
}
// ownerOnlyMiddleware wraps a handler to require Owner role (position == 100).
// It reads the user from context (set by adminAuthMiddleware) rather than
// re-authenticating, avoiding redundant DB queries and session-expiry gaps.
func ownerOnlyMiddleware(database *db.DB, next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
user, ok := r.Context().Value(adminUserKey).(*db.User)
if !ok || user == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "not authenticated")
return
}
role, err := database.GetRoleByID(user.RoleID)
if err != nil || role == nil {
writeErr(w, http.StatusForbidden, "FORBIDDEN", "role not found")
return
}
if role.Position < permissions.OwnerRolePosition {
writeErr(w, http.StatusForbidden, "FORBIDDEN", "owner role required")
return
}
next.ServeHTTP(w, r)
})
}
// ─── Helpers ──────────────────────────────────────────────────────────────────
type errorResponse struct {
Error string `json:"error"`
Message string `json:"message"`
}
func writeJSON(w http.ResponseWriter, status int, v any) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(v)
}
func writeErr(w http.ResponseWriter, status int, code, msg string) {
writeJSON(w, status, errorResponse{Error: code, Message: msg})
}
func pathInt64(r *http.Request, param string) (int64, error) {
raw := chi.URLParam(r, param)
return strconv.ParseInt(raw, 10, 64)
}
// queryInt parses an integer query parameter with a minimum and maximum bound.
// Use minVal=1 for limit parameters, minVal=0 for offset parameters.
func queryInt(r *http.Request, key string, defaultVal, minVal int) int {
raw := r.URL.Query().Get(key)
if raw == "" {
return defaultVal
}
n, err := strconv.Atoi(raw)
if err != nil || n < minVal {
return defaultVal
}
// Cap to prevent unbounded result sets exhausting memory.
const maxLimit = 500
if n > maxLimit {
return maxLimit
}
return n
}
// actorFromContext returns the authenticated user's ID stored in the request
// context by adminAuthMiddleware. Returns 0 if called outside that middleware
// (should not happen in production).
func actorFromContext(r *http.Request) int64 {
user, ok := r.Context().Value(adminUserKey).(*db.User)
if !ok || user == nil {
return 0
}
return user.ID
}
+62
View File
@@ -0,0 +1,62 @@
package admin
import (
"encoding/json"
"net/http"
"strconv"
"github.com/go-chi/chi/v5"
"github.com/owncord/server/db"
)
// ─── Helpers ──────────────────────────────────────────────────────────────────
type errorResponse struct {
Error string `json:"error"`
Message string `json:"message"`
}
func writeJSON(w http.ResponseWriter, status int, v any) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(v)
}
func writeErr(w http.ResponseWriter, status int, code, msg string) {
writeJSON(w, status, errorResponse{Error: code, Message: msg})
}
func pathInt64(r *http.Request, param string) (int64, error) {
raw := chi.URLParam(r, param)
return strconv.ParseInt(raw, 10, 64)
}
// queryInt parses an integer query parameter with a minimum and maximum bound.
// Use minVal=1 for limit parameters, minVal=0 for offset parameters.
func queryInt(r *http.Request, key string, defaultVal, minVal int) int {
raw := r.URL.Query().Get(key)
if raw == "" {
return defaultVal
}
n, err := strconv.Atoi(raw)
if err != nil || n < minVal {
return defaultVal
}
// Cap to prevent unbounded result sets exhausting memory.
const maxLimit = 500
if n > maxLimit {
return maxLimit
}
return n
}
// actorFromContext returns the authenticated user's ID stored in the request
// context by adminAuthMiddleware. Returns 0 if called outside that middleware
// (should not happen in production).
func actorFromContext(r *http.Request) int64 {
user, ok := r.Context().Value(adminUserKey).(*db.User)
if !ok || user == nil {
return 0
}
return user.ID
}
+86
View File
@@ -0,0 +1,86 @@
package admin
import (
"context"
"net/http"
"github.com/owncord/server/auth"
"github.com/owncord/server/db"
"github.com/owncord/server/permissions"
)
// ─── Middleware ───────────────────────────────────────────────────────────────
// adminAuthMiddleware validates the Bearer token and requires ADMINISTRATOR.
// On success it stores the *db.User and *db.Session in the request context so
// downstream handlers can retrieve them without re-querying the database.
func adminAuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
token, ok := auth.ExtractBearerToken(r)
if !ok {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "missing or invalid authorization header")
return
}
hash := auth.HashToken(token)
sess, err := database.GetSessionByTokenHash(hash)
if err != nil || sess == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "invalid or expired session")
return
}
if auth.IsSessionExpired(sess.ExpiresAt) {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "session has expired")
return
}
user, err := database.GetUserByID(sess.UserID)
if err != nil || user == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "user not found")
return
}
role, err := database.GetRoleByID(user.RoleID)
if err != nil || role == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "role not found")
return
}
if !permissions.HasAdmin(role.Permissions) {
writeErr(w, http.StatusForbidden, "FORBIDDEN", "administrator permission required")
return
}
ctx := context.WithValue(r.Context(), adminUserKey, user)
ctx = context.WithValue(ctx, adminSessionKey, sess)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
}
// ownerOnlyMiddleware wraps a handler to require Owner role (position == 100).
// It reads the user from context (set by adminAuthMiddleware) rather than
// re-authenticating, avoiding redundant DB queries and session-expiry gaps.
func ownerOnlyMiddleware(database *db.DB, next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
user, ok := r.Context().Value(adminUserKey).(*db.User)
if !ok || user == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "not authenticated")
return
}
role, err := database.GetRoleByID(user.RoleID)
if err != nil || role == nil {
writeErr(w, http.StatusForbidden, "FORBIDDEN", "role not found")
return
}
if role.Position < permissions.OwnerRolePosition {
writeErr(w, http.StatusForbidden, "FORBIDDEN", "owner role required")
return
}
next.ServeHTTP(w, r)
})
}
+101
View File
@@ -0,0 +1,101 @@
package admin
import "github.com/owncord/server/db"
// ─── Context keys ─────────────────────────────────────────────────────────────
// adminContextKey is an unexported type for context keys in the admin package.
type adminContextKey int
const (
// adminUserKey is the context key for the authenticated *db.User.
adminUserKey adminContextKey = iota
// adminSessionKey is the context key for the authenticated *db.Session.
adminSessionKey
)
// ─── Allowed settings keys ────────────────────────────────────────────────────
// allowedSettingKeys is the whitelist of keys that may be written via
// PATCH /admin/api/settings. Derived from the settings table in SCHEMA.md.
var allowedSettingKeys = map[string]struct{}{
"server_name": {},
"server_icon": {},
"motd": {},
"max_upload_bytes": {},
"voice_quality": {},
"require_2fa": {},
"registration_open": {},
"backup_schedule": {},
"backup_retention": {},
}
// ─── HubBroadcaster ──────────────────────────────────────────────────────────
// HubBroadcaster is the subset of ws.Hub needed by the admin package.
type HubBroadcaster interface {
BroadcastServerRestart(reason string, delaySeconds int)
BroadcastChannelCreate(ch *db.Channel)
BroadcastChannelUpdate(ch *db.Channel)
BroadcastChannelDelete(channelID int64)
BroadcastMemberBan(userID int64)
BroadcastMemberUpdate(userID int64, roleName string)
ClientCount() int
}
// ─── adminUserResponse ──────────────────────────────────────────────────────
// adminUserResponse is the safe public shape returned by user-listing and
// user-patch endpoints. It deliberately excludes PasswordHash and TOTPSecret.
type adminUserResponse struct {
ID int64 `json:"id"`
Username string `json:"username"`
Avatar *string `json:"avatar,omitempty"`
RoleID int64 `json:"role_id"`
RoleName string `json:"role_name"`
Status string `json:"status"`
CreatedAt string `json:"created_at"`
LastSeen *string `json:"last_seen,omitempty"`
Banned bool `json:"banned"`
BanReason *string `json:"ban_reason,omitempty"`
BanExpires *string `json:"ban_expires,omitempty"`
}
// toAdminUserResponse converts a db.UserWithRole to the safe response shape.
func toAdminUserResponse(u db.UserWithRole) adminUserResponse {
return adminUserResponse{
ID: u.ID,
Username: u.Username,
Avatar: u.Avatar,
RoleID: u.RoleID,
RoleName: u.RoleName,
Status: u.Status,
CreatedAt: u.CreatedAt,
LastSeen: u.LastSeen,
Banned: u.Banned,
BanReason: u.BanReason,
BanExpires: u.BanExpires,
}
}
// toAdminUserResponseFromUser converts a plain db.User to the safe response
// shape, resolving the role name via the database.
func toAdminUserResponseFromUser(database *db.DB, u *db.User) adminUserResponse {
roleName := ""
if role, err := database.GetRoleByID(u.RoleID); err == nil && role != nil {
roleName = role.Name
}
return adminUserResponse{
ID: u.ID,
Username: u.Username,
Avatar: u.Avatar,
RoleID: u.RoleID,
RoleName: roleName,
Status: u.Status,
CreatedAt: u.CreatedAt,
LastSeen: u.LastSeen,
Banned: u.Banned,
BanReason: u.BanReason,
BanExpires: u.BanExpires,
}
}
+95
View File
@@ -0,0 +1,95 @@
package permissions
import (
"errors"
"fmt"
)
// ─── Errors ─────────────────────────────────────────────────────────────────
// ErrNotDMParticipant is returned when a user is not a participant in a DM channel.
var ErrNotDMParticipant = errors.New("not a participant in this DM")
// ErrPermissionDenied is returned when a user lacks the required permission.
var ErrPermissionDenied = errors.New("permission denied")
// ─── DB interface ───────────────────────────────────────────────────────────
// ChannelOverride holds the allow/deny permission bits for a single channel.
type ChannelOverride struct {
Allow int64
Deny int64
}
// DB is the minimal database interface the Checker needs.
// Defined at the consumer (per Go convention: accept interfaces, return structs).
type DB interface {
GetChannelPermissions(channelID, roleID int64) (allow, deny int64, err error)
IsDMParticipant(userID, channelID int64) (bool, error)
}
// ─── Checker ────────────────────────────────────────────────────────────────
// Checker consolidates all channel permission checks into one reusable type.
// It is safe to share across goroutines because it holds no mutable state.
type Checker struct {
db DB
}
// NewChecker creates a Checker backed by the given database interface.
func NewChecker(db DB) *Checker {
return &Checker{db: db}
}
// HasChannelPerm reports whether the role (identified by rolePerms and roleID)
// has all the given permission bits on the specified channel. Administrator
// roles bypass all checks. Channel overrides (allow/deny) are fetched from the
// database per call.
func (ck *Checker) HasChannelPerm(rolePerms int64, roleID, channelID, perm int64) bool {
if HasAdmin(rolePerms) {
return true
}
allow, deny, err := ck.db.GetChannelPermissions(channelID, roleID)
if err != nil {
return false
}
effective := EffectivePerms(rolePerms, allow, deny)
return effective&perm == perm
}
// HasChannelPermBatch reports whether the role has the given permission on the
// channel using a pre-fetched overrides map. This avoids N+1 queries when
// filtering many channels in bulk. The zero-value ChannelOverride (no entry in
// map) is correct -- it means no override exists.
func (ck *Checker) HasChannelPermBatch(rolePerms int64, overrides map[int64]ChannelOverride, channelID, perm int64) bool {
if HasAdmin(rolePerms) {
return true
}
o := overrides[channelID] // zero-value (0, 0) when no override exists
effective := EffectivePerms(rolePerms, o.Allow, o.Deny)
return effective&perm == perm
}
// RequireChannelAccess checks whether the user can access the channel with the
// given permission. For DM channels (channelType == "dm"), it verifies
// participant membership via IsDMParticipant. For regular channels, it checks
// role-based permissions via HasChannelPerm.
//
// Returns nil on success, or a descriptive error on failure.
func (ck *Checker) RequireChannelAccess(userID, rolePerms, roleID int64, channelType string, channelID, perm int64) error {
if channelType == "dm" {
ok, err := ck.db.IsDMParticipant(userID, channelID)
if err != nil {
return fmt.Errorf("checking DM participation: %w", err)
}
if !ok {
return ErrNotDMParticipant
}
return nil
}
if !ck.HasChannelPerm(rolePerms, roleID, channelID, perm) {
return ErrPermissionDenied
}
return nil
}
+301
View File
@@ -0,0 +1,301 @@
package permissions
import (
"errors"
"testing"
)
// ─── Mock DB ────────────────────────────────────────────────────────────────
type mockDB struct {
channelPerms map[chanRoleKey]chanPerm
dmParticipants map[dmKey]bool
chanErr error
dmErr error
}
type chanRoleKey struct{ channelID, roleID int64 }
type chanPerm struct{ allow, deny int64 }
type dmKey struct{ userID, channelID int64 }
func newMockDB() *mockDB {
return &mockDB{
channelPerms: make(map[chanRoleKey]chanPerm),
dmParticipants: make(map[dmKey]bool),
}
}
func (m *mockDB) GetChannelPermissions(channelID, roleID int64) (int64, int64, error) {
if m.chanErr != nil {
return 0, 0, m.chanErr
}
key := chanRoleKey{channelID, roleID}
p, ok := m.channelPerms[key]
if !ok {
return 0, 0, nil
}
return p.allow, p.deny, nil
}
func (m *mockDB) IsDMParticipant(userID, channelID int64) (bool, error) {
if m.dmErr != nil {
return false, m.dmErr
}
return m.dmParticipants[dmKey{userID, channelID}], nil
}
// ─── HasChannelPerm tests ───────────────────────────────────────────────────
func TestHasChannelPerm(t *testing.T) {
tests := []struct {
name string
rolePerms int64
roleID int64
channelID int64
perm int64
overrides map[chanRoleKey]chanPerm
chanErr error
want bool
}{
{
name: "admin bypass returns true",
rolePerms: Administrator | SendMessages,
roleID: 1,
channelID: 10,
perm: ManageChannels,
want: true,
},
{
name: "non-admin with allow override returns true",
rolePerms: ReadMessages,
roleID: 4,
channelID: 10,
perm: SendMessages,
overrides: map[chanRoleKey]chanPerm{
{10, 4}: {allow: SendMessages, deny: 0},
},
want: true,
},
{
name: "non-admin with deny override returns false",
rolePerms: ReadMessages | SendMessages,
roleID: 4,
channelID: 10,
perm: SendMessages,
overrides: map[chanRoleKey]chanPerm{
{10, 4}: {allow: 0, deny: SendMessages},
},
want: false,
},
{
name: "non-admin without override uses base perms",
rolePerms: ReadMessages | SendMessages,
roleID: 4,
channelID: 10,
perm: SendMessages,
want: true,
},
{
name: "non-admin lacking base perm returns false",
rolePerms: ReadMessages,
roleID: 4,
channelID: 10,
perm: SendMessages,
want: false,
},
{
name: "db error returns false",
rolePerms: ReadMessages | SendMessages,
roleID: 4,
channelID: 10,
perm: SendMessages,
chanErr: errors.New("db error"),
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
db := newMockDB()
db.chanErr = tt.chanErr
for k, v := range tt.overrides {
db.channelPerms[k] = v
}
ck := NewChecker(db)
got := ck.HasChannelPerm(tt.rolePerms, tt.roleID, tt.channelID, tt.perm)
if got != tt.want {
t.Errorf("HasChannelPerm() = %v, want %v", got, tt.want)
}
})
}
}
// ─── HasChannelPermBatch tests ──────────────────────────────────────────────
func TestHasChannelPermBatch(t *testing.T) {
tests := []struct {
name string
rolePerms int64
overrides map[int64]ChannelOverride
channelID int64
perm int64
want bool
}{
{
name: "admin bypass returns true",
rolePerms: Administrator,
channelID: 10,
perm: ManageChannels,
overrides: map[int64]ChannelOverride{},
want: true,
},
{
name: "uses pre-fetched allow override",
rolePerms: ReadMessages,
channelID: 10,
perm: SendMessages,
overrides: map[int64]ChannelOverride{
10: {Allow: SendMessages, Deny: 0},
},
want: true,
},
{
name: "uses pre-fetched deny override",
rolePerms: ReadMessages | SendMessages,
channelID: 10,
perm: SendMessages,
overrides: map[int64]ChannelOverride{
10: {Allow: 0, Deny: SendMessages},
},
want: false,
},
{
name: "missing override uses base perms (zero-value)",
rolePerms: ReadMessages | SendMessages,
channelID: 99,
perm: SendMessages,
overrides: map[int64]ChannelOverride{},
want: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ck := NewChecker(newMockDB())
got := ck.HasChannelPermBatch(tt.rolePerms, tt.overrides, tt.channelID, tt.perm)
if got != tt.want {
t.Errorf("HasChannelPermBatch() = %v, want %v", got, tt.want)
}
})
}
}
// ─── RequireChannelAccess tests ─────────────────────────────────────────────
func TestRequireChannelAccess(t *testing.T) {
tests := []struct {
name string
userID int64
rolePerms int64
roleID int64
channelType string
channelID int64
perm int64
dmOK bool
dmErr error
wantErr error
}{
{
name: "DM channel - participant allowed",
userID: 1,
channelType: "dm",
channelID: 100,
dmOK: true,
wantErr: nil,
},
{
name: "DM channel - non-participant denied",
userID: 1,
channelType: "dm",
channelID: 100,
dmOK: false,
wantErr: ErrNotDMParticipant,
},
{
name: "DM channel - db error",
userID: 1,
channelType: "dm",
channelID: 100,
dmErr: errors.New("connection lost"),
},
{
name: "regular channel - has perm",
userID: 1,
rolePerms: ReadMessages | SendMessages,
roleID: 4,
channelType: "text",
channelID: 10,
perm: SendMessages,
wantErr: nil,
},
{
name: "regular channel - lacks perm",
userID: 1,
rolePerms: ReadMessages,
roleID: 4,
channelType: "text",
channelID: 10,
perm: SendMessages,
wantErr: ErrPermissionDenied,
},
{
name: "DM checks participant not role",
userID: 1,
rolePerms: 0, // no permissions at all
roleID: 0, // no role
channelType: "dm",
channelID: 100,
dmOK: true,
wantErr: nil,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
db := newMockDB()
db.dmErr = tt.dmErr
if tt.dmOK {
db.dmParticipants[dmKey{tt.userID, tt.channelID}] = true
}
ck := NewChecker(db)
err := ck.RequireChannelAccess(tt.userID, tt.rolePerms, tt.roleID, tt.channelType, tt.channelID, tt.perm)
if tt.dmErr != nil {
// Expect wrapped error.
if err == nil {
t.Fatal("RequireChannelAccess() = nil, want error")
}
if !errors.Is(err, tt.dmErr) {
t.Errorf("RequireChannelAccess() error does not wrap dmErr: got %v", err)
}
return
}
if tt.wantErr == nil {
if err != nil {
t.Errorf("RequireChannelAccess() unexpected error: %v", err)
}
return
}
if err == nil {
t.Errorf("RequireChannelAccess() = nil, want %v", tt.wantErr)
return
}
if !errors.Is(err, tt.wantErr) {
t.Errorf("RequireChannelAccess() error = %v, want %v", err, tt.wantErr)
}
})
}
}
+359
View File
@@ -0,0 +1,359 @@
// seed.go is a standalone tool that populates an OwnCord database with
// development data (users, channels, messages, DMs). It is idempotent:
// running it multiple times against the same database is safe.
//
// Usage:
//
// go run scripts/seed.go # uses ./data/chatserver.db
// go run scripts/seed.go -db path/to/owncord.db # custom path
package main
import (
"flag"
"fmt"
"log"
"os"
"github.com/owncord/server/auth"
"github.com/owncord/server/db"
)
// ─── Seed data definitions ──────────────────────────────────────────────────
// seedUser defines a user to create during seeding.
type seedUser struct {
Username string
Password string
RoleID int // 1=Owner, 2=Admin, 3=Moderator, 4=Member
}
// seedChannel defines a channel to create during seeding.
type seedChannel struct {
Name string
Type string // "text" or "voice"
Category string
Topic string
Position int
}
// seedMessage defines a message to insert during seeding.
// ChannelIdx and UserIdx refer to zero-based indices into the channels and
// users slices (resolved after creation).
type seedMessage struct {
ChannelIdx int
UserIdx int
Content string
}
var seedUsers = []seedUser{
{Username: "admin", Password: "admin123", RoleID: 1},
{Username: "alice", Password: "password123", RoleID: 4},
{Username: "bob", Password: "password123", RoleID: 4},
{Username: "charlie", Password: "password123", RoleID: 4},
}
var seedChannels = []seedChannel{
{Name: "general", Type: "text", Category: "Text Channels", Topic: "General chat", Position: 0},
{Name: "random", Type: "text", Category: "Text Channels", Topic: "Off-topic discussion", Position: 1},
{Name: "gaming", Type: "text", Category: "Text Channels", Topic: "Gaming talk", Position: 2},
{Name: "Voice Lounge", Type: "voice", Category: "Voice Channels", Position: 3},
{Name: "Gaming Voice", Type: "voice", Category: "Voice Channels", Position: 4},
}
// Channel indices for readability.
const (
chGeneral = 0
chRandom = 1
chGaming = 2
)
// User indices for readability.
const (
uAdmin = 0
uAlice = 1
uBob = 2
uCharlie = 3
)
var seedMessages = []seedMessage{
// #general
{chGeneral, uAdmin, "Welcome to OwnCord! This is the general channel."},
{chGeneral, uAlice, "Hey everyone! Glad to be here."},
{chGeneral, uBob, "Hello! This looks great."},
{chGeneral, uCharlie, "Hi all, what's everyone up to?"},
{chGeneral, uAdmin, "Feel free to chat about anything here."},
{chGeneral, uAlice, "Anyone tried the voice chat yet?"},
{chGeneral, uBob, "Not yet, but I'm about to!"},
{chGeneral, uCharlie, "The UI looks really clean."},
{chGeneral, uAdmin, "Thanks! We've been working hard on it."},
{chGeneral, uAlice, "Can we customize themes?"},
{chGeneral, uAdmin, "Yes! Check the settings panel."},
// #random
{chRandom, uBob, "Random thought: pineapple on pizza is underrated."},
{chRandom, uCharlie, "Hard disagree, but I respect your opinion."},
{chRandom, uAlice, "Let's not start a war here lol"},
{chRandom, uBob, "Too late, the war has begun!"},
{chRandom, uAdmin, "Keep it friendly, folks!"},
{chRandom, uCharlie, "Anyone watching any good shows lately?"},
{chRandom, uAlice, "I just finished a great series, highly recommend it."},
{chRandom, uBob, "What series?"},
{chRandom, uAlice, "I'll share the link later!"},
// #gaming
{chGaming, uCharlie, "Anyone up for some co-op tonight?"},
{chGaming, uBob, "I'm down! What game?"},
{chGaming, uCharlie, "Thinking something chill, maybe Minecraft?"},
{chGaming, uAlice, "Count me in!"},
{chGaming, uAdmin, "I might join later if I finish some work."},
{chGaming, uBob, "No pressure, we'll be on for a while."},
}
// seedDMMessages are messages exchanged in the admin<->alice DM channel.
var seedDMMessages = []struct {
FromIdx int // index into seedUsers
Content string
}{
{uAdmin, "Hey Alice, welcome to the server!"},
{uAlice, "Thanks! Everything looks awesome."},
{uAdmin, "Let me know if you run into any issues."},
{uAlice, "Will do! One question: how do I change my avatar?"},
{uAdmin, "Go to Settings > Account, you can upload one there."},
}
// ─── Main ───────────────────────────────────────────────────────────────────
func main() {
dbPath := flag.String("db", "data/chatserver.db", "path to the SQLite database file")
confirmDev := flag.Bool("confirm-dev", false, "confirm this is a development database (required)")
flag.Parse()
if !*confirmDev {
fmt.Fprintln(os.Stderr, "⚠ This script creates users with weak passwords.")
fmt.Fprintln(os.Stderr, " Pass -confirm-dev to confirm this is a development database.")
os.Exit(1)
}
log.SetFlags(0) // no timestamp prefix — keep output clean
database, err := db.Open(*dbPath)
if err != nil {
log.Fatalf("failed to open database at %s: %v", *dbPath, err)
}
defer database.Close()
if err := db.Migrate(database); err != nil {
log.Fatalf("failed to run migrations: %v", err)
}
userIDs, err := createUsers(database)
if err != nil {
log.Fatalf("failed to create users: %v", err)
}
channelIDs, err := createChannels(database)
if err != nil {
log.Fatalf("failed to create channels: %v", err)
}
msgCount, err := createMessages(database, channelIDs, userIDs)
if err != nil {
log.Fatalf("failed to create messages: %v", err)
}
dmMsgCount, err := createDMConversation(database, userIDs)
if err != nil {
log.Fatalf("failed to create DM conversation: %v", err)
}
fmt.Println("--- Seed complete ---")
fmt.Printf(" Users: %d\n", len(userIDs))
fmt.Printf(" Channels: %d\n", len(channelIDs))
fmt.Printf(" Messages: %d (channel) + %d (DM) = %d total\n",
msgCount, dmMsgCount, msgCount+dmMsgCount)
}
// ─── User creation ──────────────────────────────────────────────────────────
func createUsers(database *db.DB) ([]int64, error) {
ids := make([]int64, len(seedUsers))
for i, su := range seedUsers {
existing, err := database.GetUserByUsername(su.Username)
if err != nil {
return nil, fmt.Errorf("checking user %q: %w", su.Username, err)
}
if existing != nil {
ids[i] = existing.ID
fmt.Printf("[skip] user %q already exists (id=%d)\n", su.Username, existing.ID)
continue
}
hash, err := auth.HashPassword(su.Password)
if err != nil {
return nil, fmt.Errorf("hashing password for %q: %w", su.Username, err)
}
id, err := database.CreateUser(su.Username, hash, su.RoleID)
if err != nil {
return nil, fmt.Errorf("creating user %q: %w", su.Username, err)
}
ids[i] = id
roleName := roleNameFromID(su.RoleID)
fmt.Printf("[created] user %q (id=%d, role=%s)\n", su.Username, id, roleName)
}
return ids, nil
}
// roleNameFromID returns a human-readable role name for display purposes.
func roleNameFromID(roleID int) string {
switch roleID {
case 1:
return "owner"
case 2:
return "admin"
case 3:
return "moderator"
case 4:
return "member"
default:
return fmt.Sprintf("role_%d", roleID)
}
}
// ─── Channel creation ───────────────────────────────────────────────────────
func createChannels(database *db.DB) ([]int64, error) {
ids := make([]int64, len(seedChannels))
// Fetch existing channels once to check for duplicates.
existing, err := database.ListChannels()
if err != nil {
return nil, fmt.Errorf("listing channels: %w", err)
}
existingByName := make(map[string]int64, len(existing))
for _, ch := range existing {
existingByName[ch.Name] = ch.ID
}
for i, sc := range seedChannels {
if id, found := existingByName[sc.Name]; found {
ids[i] = id
fmt.Printf("[skip] channel %q already exists (id=%d)\n", sc.Name, id)
continue
}
id, err := database.CreateChannel(sc.Name, sc.Type, sc.Category, sc.Topic, sc.Position)
if err != nil {
return nil, fmt.Errorf("creating channel %q: %w", sc.Name, err)
}
ids[i] = id
fmt.Printf("[created] channel %q (id=%d, type=%s)\n", sc.Name, id, sc.Type)
}
return ids, nil
}
// ─── Message creation ───────────────────────────────────────────────────────
func createMessages(database *db.DB, channelIDs, userIDs []int64) (int, error) {
created := 0
for _, sm := range seedMessages {
channelID := channelIDs[sm.ChannelIdx]
userID := userIDs[sm.UserIdx]
// Check if this exact message already exists (content + user + channel).
exists, err := messageExists(database, channelID, userID, sm.Content)
if err != nil {
return 0, fmt.Errorf("checking message existence: %w", err)
}
if exists {
continue
}
if _, err := database.CreateMessage(channelID, userID, sm.Content, nil); err != nil {
return 0, fmt.Errorf("creating message in channel %d: %w", channelID, err)
}
created++
}
if created > 0 {
fmt.Printf("[created] %d channel messages\n", created)
} else {
fmt.Println("[skip] channel messages already seeded")
}
return created, nil
}
// messageExists checks whether a message with the given content from the given
// user already exists in the channel. Used for idempotency.
func messageExists(database *db.DB, channelID, userID int64, content string) (bool, error) {
var count int
err := database.QueryRow(
`SELECT COUNT(*) FROM messages WHERE channel_id = ? AND user_id = ? AND content = ? AND deleted = 0`,
channelID, userID, content,
).Scan(&count)
if err != nil {
return false, err
}
return count > 0, nil
}
// ─── DM conversation ────────────────────────────────────────────────────────
func createDMConversation(database *db.DB, userIDs []int64) (int, error) {
adminID := userIDs[uAdmin]
aliceID := userIDs[uAlice]
ch, isNew, err := database.GetOrCreateDMChannel(adminID, aliceID)
if err != nil {
return 0, fmt.Errorf("creating DM channel: %w", err)
}
if isNew {
fmt.Printf("[created] DM channel between admin and alice (id=%d)\n", ch.ID)
} else {
fmt.Printf("[skip] DM channel between admin and alice already exists (id=%d)\n", ch.ID)
}
created := 0
for _, dm := range seedDMMessages {
senderID := userIDs[dm.FromIdx]
exists, err := messageExists(database, ch.ID, senderID, dm.Content)
if err != nil {
return 0, fmt.Errorf("checking DM message existence: %w", err)
}
if exists {
continue
}
if _, err := database.CreateMessage(ch.ID, senderID, dm.Content, nil); err != nil {
return 0, fmt.Errorf("creating DM message: %w", err)
}
created++
}
if created > 0 {
fmt.Printf("[created] %d DM messages\n", created)
} else {
fmt.Println("[skip] DM messages already seeded")
}
return created, nil
}
// ─── Ensure data directory exists ───────────────────────────────────────────
func init() {
// The default DB path is data/chatserver.db. Ensure the data directory
// exists so db.Open doesn't fail on a fresh checkout.
if err := os.MkdirAll("data", 0o755); err != nil {
log.Printf("warning: could not create data directory: %v", err)
}
}
+3 -566
View File
@@ -9,7 +9,6 @@ import (
"github.com/microcosm-cc/bluemonday"
"github.com/owncord/server/auth"
"github.com/owncord/server/db"
"github.com/owncord/server/permissions"
)
// Rate limit windows.
@@ -42,7 +41,6 @@ func (h *Hub) HandleVoiceLeaveForTest(c *Client) {
h.handleVoiceLeave(c)
}
// handleMessage parses the envelope and dispatches to the appropriate handler.
func (h *Hub) handleMessage(c *Client, raw []byte) {
// Periodic session expiry check: every SessionCheckInterval messages,
@@ -103,522 +101,14 @@ func (h *Hub) handleMessage(c *Client, raw []byte) {
reqLog.Debug("ws ← client message")
switch env.Type {
case "chat_send":
h.handleChatSend(c, env.ID, env.Payload)
case "chat_edit":
h.handleChatEdit(c, env.ID, env.Payload)
case "chat_delete":
h.handleChatDelete(c, env.ID, env.Payload)
case "reaction_add":
h.handleReaction(c, true, env.Payload)
case "reaction_remove":
h.handleReaction(c, false, env.Payload)
case "typing_start":
h.handleTyping(c, env.Payload)
case "presence_update":
h.handlePresence(c, env.Payload)
case "channel_focus":
h.handleChannelFocus(c, env.Payload)
case "voice_join":
h.handleVoiceJoin(c, env.Payload)
case "voice_leave":
h.handleVoiceLeave(c)
case "voice_token_refresh":
h.handleVoiceTokenRefresh(c)
case "voice_mute":
h.handleVoiceMute(c, env.Payload)
case "voice_deafen":
h.handleVoiceDeafen(c, env.Payload)
case "voice_camera":
h.handleVoiceCamera(c, env.Payload)
case "voice_screenshare":
h.handleVoiceScreenshare(c, env.Payload)
case "ping":
c.sendMsg(buildJSON(map[string]any{"type": "pong"}))
default:
if !h.registry.Dispatch(env.Type, h, c, env.ID, env.Payload) {
reqLog.Warn("ws handleMessage unknown type")
c.sendMsg(buildErrorMsg(ErrCodeUnknownType, fmt.Sprintf("unknown message type: %s", env.Type)))
}
}
// handleChatSend processes a chat_send message.
func (h *Hub) handleChatSend(c *Client, reqID string, payload json.RawMessage) {
// Rate limit.
ratKey := fmt.Sprintf("chat:%d", c.userID)
if !h.limiter.Allow(ratKey, chatRateLimit, chatWindow) {
c.sendMsg(buildRateLimitError("too many messages", chatWindow.Seconds()))
return
}
var p struct {
ChannelID json.Number `json:"channel_id"`
Content string `json:"content"`
ReplyTo *int64 `json:"reply_to"`
Attachments []string `json:"attachments"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid chat_send payload"))
return
}
channelID, err := p.ChannelID.Int64()
if err != nil || channelID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "channel_id must be a positive integer"))
return
}
// Check channel exists.
ch, err := h.db.GetChannel(channelID)
if err != nil || ch == nil {
c.sendMsg(buildErrorMsg(ErrCodeNotFound, "channel not found"))
return
}
// DM channels use participant-based auth instead of role permissions.
isDM := ch.Type == "dm"
if isDM {
ok, dmErr := h.db.IsDMParticipant(c.userID, channelID)
if dmErr != nil {
slog.Error("ws handleChatSend IsDMParticipant", "err", dmErr)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to check DM participation"))
return
}
if !ok {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "you are not a participant in this DM"))
return
}
} else {
// Permission check for non-DM channels.
if !h.requireChannelPerm(c, channelID, permissions.ReadMessages|permissions.SendMessages, "SEND_MESSAGES") {
return
}
}
// Slow mode enforcement: moderators with MANAGE_MESSAGES bypass it.
// DM channels do not have slow mode.
if !isDM && ch.SlowMode > 0 && !h.hasChannelPerm(c, channelID, permissions.ManageMessages) {
slowKey := fmt.Sprintf("slow:%d:%d", c.userID, channelID)
if !h.limiter.Allow(slowKey, 1, time.Duration(ch.SlowMode)*time.Second) {
c.sendMsg(buildErrorMsg(ErrCodeSlowMode, fmt.Sprintf("channel has %ds slow mode", ch.SlowMode)))
return
}
}
// Sanitize and validate content length.
content := sanitizer.Sanitize(p.Content)
if content == "" && len(p.Attachments) == 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message content cannot be empty"))
return
}
if len([]rune(content)) > maxMessageLen {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message content exceeds maximum length of 4000 characters"))
return
}
// Check attachment permission before persisting anything.
if len(p.Attachments) > 0 {
if !h.requireChannelPerm(c, channelID, permissions.AttachFiles, "ATTACH_FILES") {
return
}
}
// Persist message.
msgID, err := h.db.CreateMessage(channelID, c.userID, content, p.ReplyTo)
if err != nil {
slog.Error("ws handleChatSend CreateMessage", "err", err)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to save message"))
return
}
// Link attachments if provided.
var attachments []map[string]any
if len(p.Attachments) > 0 {
linked, linkErr := h.db.LinkAttachmentsToMessage(msgID, p.Attachments)
if linkErr != nil {
slog.Error("ws handleChatSend LinkAttachments", "err", linkErr, "msg_id", msgID)
// Delete the orphaned message so it doesn't persist without its attachments.
if delErr := h.db.DeleteMessage(msgID, c.userID, true); delErr != nil {
slog.Error("ws handleChatSend DeleteMessage (cleanup)", "err", delErr, "msg_id", msgID)
}
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to send message with attachments"))
return
}
if linked > 0 {
attMap, attErr := h.db.GetAttachmentsByMessageIDs([]int64{msgID})
if attErr != nil {
slog.Error("ws handleChatSend GetAttachments", "err", attErr)
} else {
for _, ai := range attMap[msgID] {
attachments = append(attachments, map[string]any{
"id": ai.ID,
"filename": ai.Filename,
"size": ai.Size,
"mime": ai.Mime,
"url": ai.URL,
})
}
}
}
}
// Retrieve to get timestamp.
msg, err := h.db.GetMessage(msgID)
if err != nil || msg == nil {
slog.Error("ws handleChatSend GetMessage after create", "err", err)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to retrieve message"))
return
}
var username string
var avatar *string
if c.user != nil {
username = c.user.Username
avatar = c.user.Avatar
}
slog.Debug("message sent", "user", username, "channel_id", channelID, "msg_id", msgID)
// Ack sender.
c.sendMsg(buildChatSendOK(reqID, msgID, msg.Timestamp))
// Broadcast message.
broadcast := buildChatMessage(msgID, channelID, c.userID, username, avatar, c.roleName, content, msg.Timestamp, p.ReplyTo, attachments)
if isDM {
// DM: send directly to both participants instead of channel broadcast.
participantIDs, pErr := h.db.GetDMParticipantIDs(channelID)
if pErr != nil {
slog.Error("ws handleChatSend GetDMParticipantIDs", "err", pErr, "channel_id", channelID)
}
for _, pid := range participantIDs {
h.SendToUser(pid, broadcast)
}
// Auto-reopen the DM for the recipient if it was closed.
for _, pid := range participantIDs {
if pid == c.userID {
continue
}
if openErr := h.db.OpenDM(pid, channelID); openErr != nil {
slog.Error("ws handleChatSend OpenDM", "err", openErr,
"recipient_id", pid, "channel_id", channelID)
continue
}
// Notify the recipient that the DM was (re)opened.
// Build the event with the sender as the recipient's "other user".
if c.user != nil {
h.SendToUser(pid, buildDMChannelOpen(channelID, c.user))
}
}
} else {
h.BroadcastToChannel(channelID, broadcast)
}
}
// handleChatEdit processes a chat_edit message.
func (h *Hub) handleChatEdit(c *Client, _ string, payload json.RawMessage) {
ratKey := fmt.Sprintf("chat_edit:%d", c.userID)
if !h.limiter.Allow(ratKey, chatRateLimit, chatWindow) {
c.sendMsg(buildRateLimitError("too many edits", chatWindow.Seconds()))
return
}
var p struct {
MessageID json.Number `json:"message_id"`
Content string `json:"content"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid chat_edit payload"))
return
}
msgID, err := p.MessageID.Int64()
if err != nil || msgID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message_id must be positive integer"))
return
}
content := sanitizer.Sanitize(p.Content)
if content == "" {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "content cannot be empty"))
return
}
if len([]rune(content)) > maxMessageLen {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message too long"))
return
}
// Fetch message first to get the channel ID for the permission check.
// Use an opaque error to prevent message-ID enumeration (IDOR).
msg, err := h.db.GetMessage(msgID)
if err != nil || msg == nil {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot edit this message"))
return
}
// Check channel type for DM-aware permission handling.
editCh, chErr := h.db.GetChannel(msg.ChannelID)
editIsDM := chErr == nil && editCh != nil && editCh.Type == "dm"
if editIsDM {
ok, dmErr := h.db.IsDMParticipant(c.userID, msg.ChannelID)
if dmErr != nil || !ok {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot edit this message"))
return
}
} else {
// Re-check that the user still has SendMessages permission on this channel.
if !h.hasChannelPerm(c, msg.ChannelID, permissions.SendMessages) {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot edit this message"))
return
}
}
// EditMessage checks ownership internally.
if err := h.db.EditMessage(msgID, c.userID, content); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot edit this message"))
return
}
// Re-fetch to get the updated edited_at timestamp.
msg, err = h.db.GetMessage(msgID)
if err != nil || msg == nil {
slog.Error("ws handleChatEdit GetMessage after edit", "err", err, "msg_id", msgID)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "edit saved but broadcast failed"))
return
}
editedAt := ""
if msg.EditedAt != nil {
editedAt = *msg.EditedAt
}
slog.Debug("message edited", "user_id", c.userID, "msg_id", msgID, "channel_id", msg.ChannelID)
editedMsg := buildChatEdited(msgID, msg.ChannelID, content, editedAt)
if editIsDM {
h.broadcastToDMParticipants(msg.ChannelID, editedMsg)
} else {
h.BroadcastToChannel(msg.ChannelID, editedMsg)
}
}
// handleChatDelete processes a chat_delete message.
func (h *Hub) handleChatDelete(c *Client, _ string, payload json.RawMessage) {
ratKey := fmt.Sprintf("chat_delete:%d", c.userID)
if !h.limiter.Allow(ratKey, chatRateLimit, chatWindow) {
c.sendMsg(buildRateLimitError("too many deletes", chatWindow.Seconds()))
return
}
var p struct {
MessageID json.Number `json:"message_id"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid chat_delete payload"))
return
}
msgID, err := p.MessageID.Int64()
if err != nil || msgID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message_id must be positive integer"))
return
}
// Use an opaque error to prevent message-ID enumeration (IDOR).
msg, err := h.db.GetMessage(msgID)
if err != nil || msg == nil {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot delete this message"))
return
}
// Check channel type for DM-aware permission handling.
delCh, chErr := h.db.GetChannel(msg.ChannelID)
delIsDM := chErr == nil && delCh != nil && delCh.Type == "dm"
if delIsDM {
ok, dmErr := h.db.IsDMParticipant(c.userID, msg.ChannelID)
if dmErr != nil || !ok {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot delete this message"))
return
}
} else {
// Ensure the user still has at least ReadMessages on this channel.
if !h.hasChannelPerm(c, msg.ChannelID, permissions.ReadMessages) {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot delete this message"))
return
}
}
// In DMs, users can only delete their own messages (no mod override).
isMod := !delIsDM && h.hasChannelPerm(c, msg.ChannelID, permissions.ManageMessages)
if err := h.db.DeleteMessage(msgID, c.userID, isMod); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot delete this message"))
return
}
slog.Debug("message deleted", "user_id", c.userID, "msg_id", msgID, "channel_id", msg.ChannelID, "is_mod", isMod)
_ = h.db.LogAudit(c.userID, "message_delete", "message", msgID,
fmt.Sprintf("channel %d, mod_action=%v", msg.ChannelID, isMod))
deletedMsg := buildChatDeleted(msgID, msg.ChannelID)
if delIsDM {
h.broadcastToDMParticipants(msg.ChannelID, deletedMsg)
} else {
h.BroadcastToChannel(msg.ChannelID, deletedMsg)
}
}
// handleReaction processes reaction_add and reaction_remove messages.
func (h *Hub) handleReaction(c *Client, add bool, payload json.RawMessage) {
ratKey := fmt.Sprintf("reaction:%d", c.userID)
if !h.limiter.Allow(ratKey, reactionRateLimit, reactionWindow) {
c.sendMsg(buildRateLimitError("too many reactions", reactionWindow.Seconds()))
return
}
var p struct {
MessageID json.Number `json:"message_id"`
Emoji string `json:"emoji"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid reaction payload"))
return
}
msgID, err := p.MessageID.Int64()
if err != nil || msgID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message_id must be positive integer"))
return
}
if p.Emoji == "" {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "emoji cannot be empty"))
return
}
if len(p.Emoji) > 32 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "emoji too long"))
return
}
// Reject control characters (U+0000U+001F, U+007F) to prevent injection.
for _, r := range p.Emoji {
if r < 0x20 || r == 0x7F {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "emoji contains invalid characters"))
return
}
}
msg, err := h.db.GetMessage(msgID)
if err != nil || msg == nil {
// Normalize: return same error whether message doesn't exist or is in
// a channel the user can't see (prevents IDOR information leak).
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "reaction failed"))
return
}
// Check channel type for DM-aware permission handling.
reactCh, chErr := h.db.GetChannel(msg.ChannelID)
reactIsDM := chErr == nil && reactCh != nil && reactCh.Type == "dm"
if reactIsDM {
ok, dmErr := h.db.IsDMParticipant(c.userID, msg.ChannelID)
if dmErr != nil || !ok {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "reaction failed"))
return
}
} else {
if !h.requireChannelPerm(c, msg.ChannelID, permissions.AddReactions, "ADD_REACTIONS") {
return
}
}
action := "add"
if add {
err = h.db.AddReaction(msgID, c.userID, p.Emoji)
} else {
action = "remove"
err = h.db.RemoveReaction(msgID, c.userID, p.Emoji)
}
if err != nil {
// Sanitize: never leak raw DB constraint errors to client.
slog.Warn("reaction failed", "action", action, "msg_id", msgID, "user_id", c.userID, "err", err)
c.sendMsg(buildErrorMsg(ErrCodeConflict, "reaction failed"))
return
}
reactionMsg := buildReactionUpdate(msgID, msg.ChannelID, c.userID, p.Emoji, action)
if reactIsDM {
h.broadcastToDMParticipants(msg.ChannelID, reactionMsg)
} else {
h.BroadcastToChannel(msg.ChannelID, reactionMsg)
}
}
// handleTyping processes a typing_start message.
func (h *Hub) handleTyping(c *Client, payload json.RawMessage) {
channelID, err := parseChannelID(payload)
if err != nil || channelID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "channel_id must be positive integer"))
return
}
ratKey := fmt.Sprintf("typing:%d:%d", c.userID, channelID)
if !h.limiter.Allow(ratKey, typingRateLimit, typingWindow) {
return // silently drop; no error for typing throttle
}
// DM channels require participant check instead of role-based permissions.
typCh, typChErr := h.db.GetChannel(channelID)
if typChErr != nil || typCh == nil {
return // silently drop for unknown channels
}
if typCh.Type == "dm" {
ok, dmErr := h.db.IsDMParticipant(c.userID, channelID)
if dmErr != nil || !ok {
return // silently drop — not a DM participant
}
}
var username string
if c.user != nil {
username = c.user.Username
}
// Broadcast to channel, excluding sender.
if typCh.Type == "dm" {
h.broadcastToDMParticipantsExclude(channelID, c.userID, buildTypingMsg(channelID, c.userID, username))
} else {
h.broadcastExclude(channelID, c.userID, buildTypingMsg(channelID, c.userID, username))
}
}
// handlePresence processes a presence_update message.
func (h *Hub) handlePresence(c *Client, payload json.RawMessage) {
ratKey := fmt.Sprintf("presence:%d", c.userID)
if !h.limiter.Allow(ratKey, presenceRateLimit, presenceWindow) {
c.sendMsg(buildRateLimitError("too many presence updates", presenceWindow.Seconds()))
return
}
var p struct {
Status string `json:"status"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid presence_update payload"))
return
}
validStatuses := map[string]bool{"online": true, "idle": true, "dnd": true, "offline": true}
if !validStatuses[p.Status] {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "status must be online|idle|dnd|offline"))
return
}
if err := h.db.UpdateUserStatus(c.userID, p.Status); err != nil {
slog.Error("ws handlePresence UpdateUserStatus", "err", err, "user_id", c.userID)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to update status"))
return
}
h.BroadcastToAll(buildPresenceMsg(c.userID, p.Status))
}
// hasChannelPerm reports whether the client's role has all the given permission bits.
// The ADMINISTRATOR bit bypasses all checks.
// Delegates to the unified permissions.Checker.
func (h *Hub) hasChannelPerm(c *Client, channelID int64, perm int64) bool {
if c.user == nil {
return false
@@ -627,16 +117,7 @@ func (h *Hub) hasChannelPerm(c *Client, channelID int64, perm int64) bool {
if err != nil || role == nil {
return false
}
if role.Permissions&permissions.Administrator != 0 {
return true
}
// Check channel overrides.
allow, deny, err := h.db.GetChannelPermissions(channelID, role.ID)
if err != nil {
return false
}
effective := permissions.EffectivePerms(role.Permissions, allow, deny)
return effective&perm == perm
return h.permChecker.HasChannelPerm(role.Permissions, role.ID, channelID, perm)
}
// requireChannelPerm checks whether the client has the given permission on the
@@ -701,47 +182,3 @@ func (h *Hub) broadcastToDMParticipantsExclude(channelID, excludeUserID int64, m
h.SendToUser(pid, msg)
}
}
// handleChannelFocus sets which channel the client is currently viewing,
// so channel-scoped broadcasts (chat messages, typing) reach them.
// Also updates read_states so unread counts decrease when the user views a channel.
func (h *Hub) handleChannelFocus(c *Client, payload json.RawMessage) {
chID, err := parseChannelID(payload)
if err != nil || chID <= 0 {
slog.Debug("handleChannelFocus: invalid channel_id", "user_id", c.userID, "err", err)
return
}
// DM channels use participant-based auth instead of role-based permissions.
ch, chErr := h.db.GetChannel(chID)
if chErr != nil || ch == nil {
slog.Debug("handleChannelFocus: channel not found", "channel_id", chID)
return
}
if ch.Type == "dm" {
ok, dmErr := h.db.IsDMParticipant(c.userID, chID)
if dmErr != nil || !ok {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "not a participant in this DM"))
return
}
} else {
if !h.requireChannelPerm(c, chID, permissions.ReadMessages, "READ_MESSAGES") {
return
}
}
c.mu.Lock()
prevCh := c.channelID
c.channelID = chID
c.mu.Unlock()
slog.Debug("channel_focus", "user_id", c.userID, "channel_id", chID, "prev_channel_id", prevCh)
// Mark channel as read by updating read_states to the latest message.
latestID, latestErr := h.db.GetLatestMessageID(chID)
if latestErr == nil && latestID > 0 {
if rsErr := h.db.UpdateReadState(c.userID, chID, latestID); rsErr != nil {
slog.Warn("handleChannelFocus UpdateReadState", "err", rsErr, "user_id", c.userID, "channel_id", chID)
}
}
}
+348
View File
@@ -0,0 +1,348 @@
package ws
import (
"encoding/json"
"fmt"
"log/slog"
"time"
"github.com/owncord/server/permissions"
)
// registerChatHandlers registers all chat-related message handlers.
func registerChatHandlers(r *HandlerRegistry) {
r.Register(MsgTypeChatSend, func(h *Hub, c *Client, reqID string, payload json.RawMessage) {
h.handleChatSend(c, reqID, payload)
})
r.Register(MsgTypeChatEdit, func(h *Hub, c *Client, reqID string, payload json.RawMessage) {
h.handleChatEdit(c, reqID, payload)
})
r.Register(MsgTypeChatDelete, func(h *Hub, c *Client, reqID string, payload json.RawMessage) {
h.handleChatDelete(c, reqID, payload)
})
}
// handleChatSend processes a chat_send message.
func (h *Hub) handleChatSend(c *Client, reqID string, payload json.RawMessage) {
// Rate limit.
ratKey := fmt.Sprintf("chat:%d", c.userID)
if !h.limiter.Allow(ratKey, chatRateLimit, chatWindow) {
c.sendMsg(buildRateLimitError("too many messages", chatWindow.Seconds()))
return
}
var p struct {
ChannelID json.Number `json:"channel_id"`
Content string `json:"content"`
ReplyTo *int64 `json:"reply_to"`
Attachments []string `json:"attachments"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid chat_send payload"))
return
}
channelID, err := p.ChannelID.Int64()
if err != nil || channelID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "channel_id must be a positive integer"))
return
}
// Check channel exists.
ch, err := h.db.GetChannel(channelID)
if err != nil || ch == nil {
c.sendMsg(buildErrorMsg(ErrCodeNotFound, "channel not found"))
return
}
// DM channels use participant-based auth instead of role permissions.
isDM := ch.Type == "dm"
if isDM {
ok, dmErr := h.db.IsDMParticipant(c.userID, channelID)
if dmErr != nil {
slog.Error("ws handleChatSend IsDMParticipant", "err", dmErr)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to check DM participation"))
return
}
if !ok {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "you are not a participant in this DM"))
return
}
} else {
// Permission check for non-DM channels.
if !h.requireChannelPerm(c, channelID, permissions.ReadMessages|permissions.SendMessages, "SEND_MESSAGES") {
return
}
}
// Slow mode enforcement: moderators with MANAGE_MESSAGES bypass it.
// DM channels do not have slow mode.
if !isDM && ch.SlowMode > 0 && !h.hasChannelPerm(c, channelID, permissions.ManageMessages) {
slowKey := fmt.Sprintf("slow:%d:%d", c.userID, channelID)
if !h.limiter.Allow(slowKey, 1, time.Duration(ch.SlowMode)*time.Second) {
c.sendMsg(buildErrorMsg(ErrCodeSlowMode, fmt.Sprintf("channel has %ds slow mode", ch.SlowMode)))
return
}
}
// Sanitize and validate content length.
content := sanitizer.Sanitize(p.Content)
if content == "" && len(p.Attachments) == 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message content cannot be empty"))
return
}
if len([]rune(content)) > maxMessageLen {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message content exceeds maximum length of 4000 characters"))
return
}
// Check attachment permission before persisting anything.
// DM channels use participant-based auth (already checked above), not role permissions.
if !isDM && len(p.Attachments) > 0 {
if !h.requireChannelPerm(c, channelID, permissions.AttachFiles, "ATTACH_FILES") {
return
}
}
// Persist message.
msgID, err := h.db.CreateMessage(channelID, c.userID, content, p.ReplyTo)
if err != nil {
slog.Error("ws handleChatSend CreateMessage", "err", err)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to save message"))
return
}
// Link attachments if provided.
var attachments []map[string]any
if len(p.Attachments) > 0 {
linked, linkErr := h.db.LinkAttachmentsToMessage(msgID, p.Attachments)
if linkErr != nil {
slog.Error("ws handleChatSend LinkAttachments", "err", linkErr, "msg_id", msgID)
// Delete the orphaned message so it doesn't persist without its attachments.
if delErr := h.db.DeleteMessage(msgID, c.userID, true); delErr != nil {
slog.Error("ws handleChatSend DeleteMessage (cleanup)", "err", delErr, "msg_id", msgID)
}
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to send message with attachments"))
return
}
if linked > 0 {
attMap, attErr := h.db.GetAttachmentsByMessageIDs([]int64{msgID})
if attErr != nil {
slog.Error("ws handleChatSend GetAttachments", "err", attErr)
} else {
for _, ai := range attMap[msgID] {
attachments = append(attachments, map[string]any{
"id": ai.ID,
"filename": ai.Filename,
"size": ai.Size,
"mime": ai.Mime,
"url": ai.URL,
})
}
}
}
}
// Retrieve to get timestamp.
msg, err := h.db.GetMessage(msgID)
if err != nil || msg == nil {
slog.Error("ws handleChatSend GetMessage after create", "err", err)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to retrieve message"))
return
}
var username string
var avatar *string
if c.user != nil {
username = c.user.Username
avatar = c.user.Avatar
}
slog.Debug("message sent", "user", username, "channel_id", channelID, "msg_id", msgID)
// Ack sender.
c.sendMsg(buildChatSendOK(reqID, msgID, msg.Timestamp))
// Broadcast message.
broadcast := buildChatMessage(msgID, channelID, c.userID, username, avatar, c.roleName, content, msg.Timestamp, p.ReplyTo, attachments)
if isDM {
// DM: send directly to both participants instead of channel broadcast.
participantIDs, pErr := h.db.GetDMParticipantIDs(channelID)
if pErr != nil {
slog.Error("ws handleChatSend GetDMParticipantIDs", "err", pErr, "channel_id", channelID)
}
for _, pid := range participantIDs {
h.SendToUser(pid, broadcast)
}
// Auto-reopen the DM for the recipient if it was closed.
for _, pid := range participantIDs {
if pid == c.userID {
continue
}
if openErr := h.db.OpenDM(pid, channelID); openErr != nil {
slog.Error("ws handleChatSend OpenDM", "err", openErr,
"recipient_id", pid, "channel_id", channelID)
continue
}
// Notify the recipient that the DM was (re)opened.
// Build the event with the sender as the recipient's "other user".
if c.user != nil {
h.SendToUser(pid, buildDMChannelOpen(channelID, c.user))
}
}
} else {
h.BroadcastToChannel(channelID, broadcast)
}
}
// handleChatEdit processes a chat_edit message.
func (h *Hub) handleChatEdit(c *Client, _ string, payload json.RawMessage) {
ratKey := fmt.Sprintf("chat_edit:%d", c.userID)
if !h.limiter.Allow(ratKey, chatRateLimit, chatWindow) {
c.sendMsg(buildRateLimitError("too many edits", chatWindow.Seconds()))
return
}
var p struct {
MessageID json.Number `json:"message_id"`
Content string `json:"content"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid chat_edit payload"))
return
}
msgID, err := p.MessageID.Int64()
if err != nil || msgID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message_id must be positive integer"))
return
}
content := sanitizer.Sanitize(p.Content)
if content == "" {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "content cannot be empty"))
return
}
if len([]rune(content)) > maxMessageLen {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message too long"))
return
}
// Fetch message first to get the channel ID for the permission check.
// Use an opaque error to prevent message-ID enumeration (IDOR).
msg, err := h.db.GetMessage(msgID)
if err != nil || msg == nil {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot edit this message"))
return
}
// Check channel type for DM-aware permission handling.
editCh, chErr := h.db.GetChannel(msg.ChannelID)
editIsDM := chErr == nil && editCh != nil && editCh.Type == "dm"
if editIsDM {
ok, dmErr := h.db.IsDMParticipant(c.userID, msg.ChannelID)
if dmErr != nil || !ok {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot edit this message"))
return
}
} else {
// Re-check that the user still has SendMessages permission on this channel.
if !h.hasChannelPerm(c, msg.ChannelID, permissions.SendMessages) {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot edit this message"))
return
}
}
// EditMessage checks ownership internally.
if err := h.db.EditMessage(msgID, c.userID, content); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot edit this message"))
return
}
// Re-fetch to get the updated edited_at timestamp.
msg, err = h.db.GetMessage(msgID)
if err != nil || msg == nil {
slog.Error("ws handleChatEdit GetMessage after edit", "err", err, "msg_id", msgID)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "edit saved but broadcast failed"))
return
}
editedAt := ""
if msg.EditedAt != nil {
editedAt = *msg.EditedAt
}
slog.Debug("message edited", "user_id", c.userID, "msg_id", msgID, "channel_id", msg.ChannelID)
editedMsg := buildChatEdited(msgID, msg.ChannelID, content, editedAt)
if editIsDM {
h.broadcastToDMParticipants(msg.ChannelID, editedMsg)
} else {
h.BroadcastToChannel(msg.ChannelID, editedMsg)
}
}
// handleChatDelete processes a chat_delete message.
func (h *Hub) handleChatDelete(c *Client, _ string, payload json.RawMessage) {
ratKey := fmt.Sprintf("chat_delete:%d", c.userID)
if !h.limiter.Allow(ratKey, chatRateLimit, chatWindow) {
c.sendMsg(buildRateLimitError("too many deletes", chatWindow.Seconds()))
return
}
var p struct {
MessageID json.Number `json:"message_id"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid chat_delete payload"))
return
}
msgID, err := p.MessageID.Int64()
if err != nil || msgID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message_id must be positive integer"))
return
}
// Use an opaque error to prevent message-ID enumeration (IDOR).
msg, err := h.db.GetMessage(msgID)
if err != nil || msg == nil {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot delete this message"))
return
}
// Check channel type for DM-aware permission handling.
delCh, chErr := h.db.GetChannel(msg.ChannelID)
delIsDM := chErr == nil && delCh != nil && delCh.Type == "dm"
if delIsDM {
ok, dmErr := h.db.IsDMParticipant(c.userID, msg.ChannelID)
if dmErr != nil || !ok {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot delete this message"))
return
}
} else {
// Ensure the user still has at least ReadMessages on this channel.
if !h.hasChannelPerm(c, msg.ChannelID, permissions.ReadMessages) {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot delete this message"))
return
}
}
// In DMs, users can only delete their own messages (no mod override).
isMod := !delIsDM && h.hasChannelPerm(c, msg.ChannelID, permissions.ManageMessages)
if err := h.db.DeleteMessage(msgID, c.userID, isMod); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "cannot delete this message"))
return
}
slog.Debug("message deleted", "user_id", c.userID, "msg_id", msgID, "channel_id", msg.ChannelID, "is_mod", isMod)
_ = h.db.LogAudit(c.userID, "message_delete", "message", msgID,
fmt.Sprintf("channel %d, mod_action=%v", msg.ChannelID, isMod))
deletedMsg := buildChatDeleted(msgID, msg.ChannelID)
if delIsDM {
h.broadcastToDMParticipants(msg.ChannelID, deletedMsg)
} else {
h.BroadcastToChannel(msg.ChannelID, deletedMsg)
}
}
+10
View File
@@ -0,0 +1,10 @@
package ws
import "encoding/json"
// registerPingHandler registers the ping/pong handler.
func registerPingHandler(r *HandlerRegistry) {
r.Register(MsgTypePing, func(_ *Hub, c *Client, _ string, _ json.RawMessage) {
c.sendMsg(buildJSON(map[string]any{"type": MsgTypePong}))
})
}
+138
View File
@@ -0,0 +1,138 @@
package ws
import (
"encoding/json"
"fmt"
"log/slog"
"github.com/owncord/server/permissions"
)
// registerPresenceHandlers registers presence, typing, and channel focus handlers.
func registerPresenceHandlers(r *HandlerRegistry) {
r.Register(MsgTypeTypingStart, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handleTyping(c, payload)
})
r.Register(MsgTypePresenceUpdate, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handlePresence(c, payload)
})
r.Register(MsgTypeChannelFocus, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handleChannelFocus(c, payload)
})
}
// handleTyping processes a typing_start message.
func (h *Hub) handleTyping(c *Client, payload json.RawMessage) {
channelID, err := parseChannelID(payload)
if err != nil || channelID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "channel_id must be positive integer"))
return
}
ratKey := fmt.Sprintf("typing:%d:%d", c.userID, channelID)
if !h.limiter.Allow(ratKey, typingRateLimit, typingWindow) {
return // silently drop; no error for typing throttle
}
// DM channels require participant check instead of role-based permissions.
typCh, typChErr := h.db.GetChannel(channelID)
if typChErr != nil || typCh == nil {
return // silently drop for unknown channels
}
if typCh.Type == "dm" {
ok, dmErr := h.db.IsDMParticipant(c.userID, channelID)
if dmErr != nil || !ok {
return // silently drop — not a DM participant
}
} else {
if !h.hasChannelPerm(c, channelID, permissions.ReadMessages) {
return // silently drop — no read permission on this channel
}
}
var username string
if c.user != nil {
username = c.user.Username
}
// Broadcast to channel, excluding sender.
if typCh.Type == "dm" {
h.broadcastToDMParticipantsExclude(channelID, c.userID, buildTypingMsg(channelID, c.userID, username))
} else {
h.broadcastExclude(channelID, c.userID, buildTypingMsg(channelID, c.userID, username))
}
}
// handlePresence processes a presence_update message.
func (h *Hub) handlePresence(c *Client, payload json.RawMessage) {
ratKey := fmt.Sprintf("presence:%d", c.userID)
if !h.limiter.Allow(ratKey, presenceRateLimit, presenceWindow) {
c.sendMsg(buildRateLimitError("too many presence updates", presenceWindow.Seconds()))
return
}
var p struct {
Status string `json:"status"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid presence_update payload"))
return
}
validStatuses := map[string]bool{"online": true, "idle": true, "dnd": true, "offline": true}
if !validStatuses[p.Status] {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "status must be online|idle|dnd|offline"))
return
}
if err := h.db.UpdateUserStatus(c.userID, p.Status); err != nil {
slog.Error("ws handlePresence UpdateUserStatus", "err", err, "user_id", c.userID)
c.sendMsg(buildErrorMsg(ErrCodeInternal, "failed to update status"))
return
}
h.BroadcastToAll(buildPresenceMsg(c.userID, p.Status))
}
// handleChannelFocus sets which channel the client is currently viewing,
// so channel-scoped broadcasts (chat messages, typing) reach them.
// Also updates read_states so unread counts decrease when the user views a channel.
func (h *Hub) handleChannelFocus(c *Client, payload json.RawMessage) {
chID, err := parseChannelID(payload)
if err != nil || chID <= 0 {
slog.Debug("handleChannelFocus: invalid channel_id", "user_id", c.userID, "err", err)
return
}
// DM channels use participant-based auth instead of role-based permissions.
ch, chErr := h.db.GetChannel(chID)
if chErr != nil || ch == nil {
slog.Debug("handleChannelFocus: channel not found", "channel_id", chID)
return
}
if ch.Type == "dm" {
ok, dmErr := h.db.IsDMParticipant(c.userID, chID)
if dmErr != nil || !ok {
c.sendMsg(buildErrorMsg(ErrCodeForbidden, "not a participant in this DM"))
return
}
} else {
if !h.requireChannelPerm(c, chID, permissions.ReadMessages, "READ_MESSAGES") {
return
}
}
c.mu.Lock()
prevCh := c.channelID
c.channelID = chID
c.mu.Unlock()
slog.Debug("channel_focus", "user_id", c.userID, "channel_id", chID, "prev_channel_id", prevCh)
// Mark channel as read by updating read_states to the latest message.
latestID, latestErr := h.db.GetLatestMessageID(chID)
if latestErr == nil && latestID > 0 {
if rsErr := h.db.UpdateReadState(c.userID, chID, latestID); rsErr != nil {
slog.Warn("handleChannelFocus UpdateReadState", "err", rsErr, "user_id", c.userID, "channel_id", chID)
}
}
}
+102
View File
@@ -0,0 +1,102 @@
package ws
import (
"encoding/json"
"fmt"
"log/slog"
"github.com/owncord/server/permissions"
)
// registerReactionHandlers registers reaction_add and reaction_remove handlers.
func registerReactionHandlers(r *HandlerRegistry) {
r.Register(MsgTypeReactionAdd, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handleReaction(c, true, payload)
})
r.Register(MsgTypeReactionRemove, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handleReaction(c, false, payload)
})
}
// handleReaction processes reaction_add and reaction_remove messages.
func (h *Hub) handleReaction(c *Client, add bool, payload json.RawMessage) {
ratKey := fmt.Sprintf("reaction:%d", c.userID)
if !h.limiter.Allow(ratKey, reactionRateLimit, reactionWindow) {
c.sendMsg(buildRateLimitError("too many reactions", reactionWindow.Seconds()))
return
}
var p struct {
MessageID json.Number `json:"message_id"`
Emoji string `json:"emoji"`
}
if err := json.Unmarshal(payload, &p); err != nil {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "invalid reaction payload"))
return
}
msgID, err := p.MessageID.Int64()
if err != nil || msgID <= 0 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "message_id must be positive integer"))
return
}
if p.Emoji == "" {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "emoji cannot be empty"))
return
}
if len(p.Emoji) > 32 {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "emoji too long"))
return
}
// Reject control characters (U+0000-U+001F, U+007F) to prevent injection.
for _, r := range p.Emoji {
if r < 0x20 || r == 0x7F {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "emoji contains invalid characters"))
return
}
}
msg, err := h.db.GetMessage(msgID)
if err != nil || msg == nil {
// Normalize: return same error whether message doesn't exist or is in
// a channel the user can't see (prevents IDOR information leak).
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "reaction failed"))
return
}
// Check channel type for DM-aware permission handling.
reactCh, chErr := h.db.GetChannel(msg.ChannelID)
reactIsDM := chErr == nil && reactCh != nil && reactCh.Type == "dm"
if reactIsDM {
ok, dmErr := h.db.IsDMParticipant(c.userID, msg.ChannelID)
if dmErr != nil || !ok {
c.sendMsg(buildErrorMsg(ErrCodeBadRequest, "reaction failed"))
return
}
} else {
if !h.requireChannelPerm(c, msg.ChannelID, permissions.AddReactions, "ADD_REACTIONS") {
return
}
}
action := "add"
if add {
err = h.db.AddReaction(msgID, c.userID, p.Emoji)
} else {
action = "remove"
err = h.db.RemoveReaction(msgID, c.userID, p.Emoji)
}
if err != nil {
// Sanitize: never leak raw DB constraint errors to client.
slog.Warn("reaction failed", "action", action, "msg_id", msgID, "user_id", c.userID, "err", err)
c.sendMsg(buildErrorMsg(ErrCodeConflict, "reaction failed"))
return
}
reactionMsg := buildReactionUpdate(msgID, msg.ChannelID, c.userID, p.Emoji, action)
if reactIsDM {
h.broadcastToDMParticipants(msg.ChannelID, reactionMsg)
} else {
h.BroadcastToChannel(msg.ChannelID, reactionMsg)
}
}
+31
View File
@@ -0,0 +1,31 @@
package ws
import "encoding/json"
// registerVoiceHandlers registers all voice-related message handlers.
// The handler methods themselves live in voice_join.go, voice_leave.go,
// voice_controls.go, and voice_broadcast.go — this function only wires
// them into the registry.
func registerVoiceHandlers(r *HandlerRegistry) {
r.Register(MsgTypeVoiceJoin, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handleVoiceJoin(c, payload)
})
r.Register(MsgTypeVoiceLeave, func(h *Hub, c *Client, _ string, _ json.RawMessage) {
h.handleVoiceLeave(c)
})
r.Register(MsgTypeVoiceTokenRefresh, func(h *Hub, c *Client, _ string, _ json.RawMessage) {
h.handleVoiceTokenRefresh(c)
})
r.Register(MsgTypeVoiceMute, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handleVoiceMute(c, payload)
})
r.Register(MsgTypeVoiceDeafen, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handleVoiceDeafen(c, payload)
})
r.Register(MsgTypeVoiceCamera, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handleVoiceCamera(c, payload)
})
r.Register(MsgTypeVoiceScreenshare, func(h *Hub, c *Client, _ string, payload json.RawMessage) {
h.handleVoiceScreenshare(c, payload)
})
}
+12
View File
@@ -11,6 +11,7 @@ import (
"github.com/owncord/server/auth"
"github.com/owncord/server/db"
"github.com/owncord/server/permissions"
)
// broadcastMsg is an internal message queued for delivery.
@@ -33,6 +34,8 @@ type Hub struct {
stopOnce sync.Once
livekit *LiveKitClient
lkProcess *LiveKitProcess
registry *HandlerRegistry
permChecker *permissions.Checker
seq uint64 // atomic monotonic sequence counter
replayBuf *EventRingBuffer // recent broadcast events for reconnection replay
@@ -47,6 +50,13 @@ type Hub struct {
// NewHub creates a Hub ready to be started with Run.
// It also initializes the settings cache from the database.
func NewHub(database *db.DB, limiter *auth.RateLimiter) *Hub {
reg := NewHandlerRegistry()
registerChatHandlers(reg)
registerPresenceHandlers(reg)
registerReactionHandlers(reg)
registerVoiceHandlers(reg)
registerPingHandler(reg)
h := &Hub{
clients: make(map[int64]*Client),
db: database,
@@ -56,6 +66,8 @@ func NewHub(database *db.DB, limiter *auth.RateLimiter) *Hub {
unregister: make(chan *Client, 32),
stop: make(chan struct{}),
replayBuf: NewEventRingBuffer(1000),
registry: reg,
permChecker: permissions.NewChecker(database),
settingsName: "OwnCord Server",
settingsMotd: "Welcome!",
}
+57
View File
@@ -0,0 +1,57 @@
package ws
// WebSocket protocol message type constants.
// Generated from docs/protocol-schema.json — single source of truth for
// both Server (Go) and Client (TypeScript).
//
// Client → Server message types (received by handlers).
const (
MsgTypeAuth = "auth"
MsgTypeChatSend = "chat_send"
MsgTypeChatEdit = "chat_edit"
MsgTypeChatDelete = "chat_delete"
MsgTypeReactionAdd = "reaction_add"
MsgTypeReactionRemove = "reaction_remove"
MsgTypeTypingStart = "typing_start"
MsgTypeChannelFocus = "channel_focus"
MsgTypePresenceUpdate = "presence_update"
MsgTypeVoiceJoin = "voice_join"
MsgTypeVoiceLeave = "voice_leave"
MsgTypeVoiceMute = "voice_mute"
MsgTypeVoiceDeafen = "voice_deafen"
MsgTypeVoiceCamera = "voice_camera"
MsgTypeVoiceScreenshare = "voice_screenshare"
MsgTypePing = "ping"
MsgTypeVoiceTokenRefresh = "voice_token_refresh"
)
// Server → Client message types (sent in broadcasts/responses).
const (
MsgTypeAuthOK = "auth_ok"
MsgTypeAuthError = "auth_error"
MsgTypeReady = "ready"
MsgTypeChatMessage = "chat_message"
MsgTypeChatSendOK = "chat_send_ok"
MsgTypeChatEdited = "chat_edited"
MsgTypeChatDeleted = "chat_deleted"
MsgTypeReactionUpdate = "reaction_update"
MsgTypeTyping = "typing"
MsgTypePresence = "presence"
MsgTypeChannelCreate = "channel_create"
MsgTypeChannelUpdate = "channel_update"
MsgTypeChannelDelete = "channel_delete"
MsgTypeVoiceState = "voice_state"
MsgTypeVoiceConfig = "voice_config"
MsgTypeVoiceToken = "voice_token"
MsgTypeVoiceSpeakers = "voice_speakers"
MsgTypeVoiceLeaveBC = "voice_leave" // broadcast (same string as client msg)
MsgTypeMemberJoin = "member_join"
MsgTypeMemberLeave = "member_leave"
MsgTypeMemberUpdate = "member_update"
MsgTypeMemberBan = "member_ban"
MsgTypeServerRestart = "server_restart"
MsgTypeError = "error"
MsgTypePong = "pong"
MsgTypeDMChannelOpen = "dm_channel_open"
MsgTypeDMChannelClose = "dm_channel_close"
)
+22 -22
View File
@@ -177,7 +177,7 @@ func buildJSON(v any) []byte {
// buildErrorMsg produces an error envelope with the given code and message.
func buildErrorMsg(code, message string) []byte {
return buildJSON(map[string]any{
"type": "error",
"type": MsgTypeError,
"payload": map[string]string{
"code": code,
"message": message,
@@ -188,7 +188,7 @@ func buildErrorMsg(code, message string) []byte {
// buildRateLimitError produces a RATE_LIMITED error with retry_after per PROTOCOL.md.
func buildRateLimitError(message string, retryAfterSeconds float64) []byte {
return buildJSON(map[string]any{
"type": "error",
"type": MsgTypeError,
"payload": map[string]any{
"code": "RATE_LIMITED",
"message": message,
@@ -201,7 +201,7 @@ func buildRateLimitError(message string, retryAfterSeconds float64) []byte {
// The client treats this type as non-recoverable and stops reconnecting.
func buildAuthError(message string) []byte {
return buildJSON(map[string]any{
"type": "auth_error",
"type": MsgTypeAuthError,
"payload": map[string]string{
"message": message,
},
@@ -215,7 +215,7 @@ func buildAuthError(message string) []byte {
// buildPresenceMsg constructs a presence broadcast payload.
func buildPresenceMsg(userID int64, status string) []byte {
return buildJSON(wsMsg{
Type: "presence",
Type: MsgTypePresence,
Payload: presencePayload{UserID: userID, Status: status},
})
}
@@ -223,7 +223,7 @@ func buildPresenceMsg(userID int64, status string) []byte {
// buildMemberJoin constructs a member_join broadcast for when a user comes online.
func buildMemberJoin(user *db.User, roleName string) []byte {
return buildJSON(wsMsg{
Type: "member_join",
Type: MsgTypeMemberJoin,
Payload: memberJoinPayload{
User: memberUserPayload{
ID: user.ID,
@@ -242,7 +242,7 @@ func buildChatMessage(msgID, channelID, userID int64, username string, avatar *s
attachments = []map[string]any{}
}
return buildJSON(wsMsg{
Type: "chat_message",
Type: MsgTypeChatMessage,
Payload: chatMessagePayload{
ID: msgID,
ChannelID: channelID,
@@ -265,7 +265,7 @@ func buildChatMessage(msgID, channelID, userID int64, username string, avatar *s
// buildMemberUpdate constructs a member_update broadcast per PROTOCOL.md.
func buildMemberUpdate(userID int64, roleName string) []byte {
return buildJSON(wsMsg{
Type: "member_update",
Type: MsgTypeMemberUpdate,
Payload: memberUpdatePayload{UserID: userID, Role: roleName},
})
}
@@ -273,7 +273,7 @@ func buildMemberUpdate(userID int64, roleName string) []byte {
// buildMemberBan constructs a member_ban broadcast per PROTOCOL.md.
func buildMemberBan(userID int64) []byte {
return buildJSON(wsMsg{
Type: "member_ban",
Type: MsgTypeMemberBan,
Payload: memberBanPayload{UserID: userID},
})
}
@@ -281,7 +281,7 @@ func buildMemberBan(userID int64) []byte {
// buildChatSendOK constructs a chat_send_ok ack.
func buildChatSendOK(requestID string, msgID int64, timestamp string) []byte {
return buildJSON(wsMsg{
Type: "chat_send_ok",
Type: MsgTypeChatSendOK,
ID: requestID,
Payload: chatSendOKPayload{MessageID: msgID, Timestamp: timestamp},
})
@@ -290,7 +290,7 @@ func buildChatSendOK(requestID string, msgID int64, timestamp string) []byte {
// buildChatEdited constructs a chat_edited broadcast.
func buildChatEdited(msgID, channelID int64, content, editedAt string) []byte {
return buildJSON(wsMsg{
Type: "chat_edited",
Type: MsgTypeChatEdited,
Payload: chatEditedPayload{
MessageID: msgID,
ChannelID: channelID,
@@ -303,7 +303,7 @@ func buildChatEdited(msgID, channelID int64, content, editedAt string) []byte {
// buildChatDeleted constructs a chat_deleted broadcast.
func buildChatDeleted(msgID, channelID int64) []byte {
return buildJSON(wsMsg{
Type: "chat_deleted",
Type: MsgTypeChatDeleted,
Payload: chatDeletedPayload{MessageID: msgID, ChannelID: channelID},
})
}
@@ -311,7 +311,7 @@ func buildChatDeleted(msgID, channelID int64) []byte {
// buildReactionUpdate constructs a reaction_update broadcast.
func buildReactionUpdate(msgID, channelID, userID int64, emoji, action string) []byte {
return buildJSON(wsMsg{
Type: "reaction_update",
Type: MsgTypeReactionUpdate,
Payload: reactionUpdatePayload{
MessageID: msgID,
ChannelID: channelID,
@@ -325,7 +325,7 @@ func buildReactionUpdate(msgID, channelID, userID int64, emoji, action string) [
// buildTypingMsg constructs a typing broadcast.
func buildTypingMsg(channelID, userID int64, username string) []byte {
return buildJSON(wsMsg{
Type: "typing",
Type: MsgTypeTyping,
Payload: typingPayload{
ChannelID: channelID,
UserID: userID,
@@ -337,7 +337,7 @@ func buildTypingMsg(channelID, userID int64, username string) []byte {
// buildVoiceState constructs a voice_state server->client broadcast.
func buildVoiceState(state db.VoiceState) []byte {
return buildJSON(wsMsg{
Type: "voice_state",
Type: MsgTypeVoiceState,
Payload: voiceStatePayload{
ChannelID: state.ChannelID,
UserID: state.UserID,
@@ -354,7 +354,7 @@ func buildVoiceState(state db.VoiceState) []byte {
// buildVoiceConfig constructs a voice_config message sent after voice_join acceptance.
func buildVoiceConfig(channelID int64, quality string, bitrate int, maxUsers int) []byte {
return buildJSON(wsMsg{
Type: "voice_config",
Type: MsgTypeVoiceConfig,
Payload: voiceConfigPayload{
ChannelID: channelID,
Quality: quality,
@@ -372,7 +372,7 @@ func buildVoiceConfig(channelID int64, quality string, bitrate int, maxUsers int
// LiveKit URL (e.g. "ws://localhost:7880") for localhost clients.
func buildVoiceToken(channelID int64, token string, proxyPath string, directURL string) []byte {
return buildJSON(wsMsg{
Type: "voice_token",
Type: MsgTypeVoiceToken,
Payload: voiceTokenPayload{
ChannelID: channelID,
Token: token,
@@ -385,7 +385,7 @@ func buildVoiceToken(channelID int64, token string, proxyPath string, directURL
// buildVoiceLeave constructs a voice_leave server->client broadcast.
func buildVoiceLeave(channelID, userID int64) []byte {
return buildJSON(wsMsg{
Type: "voice_leave",
Type: MsgTypeVoiceLeaveBC,
Payload: voiceLeavePayload{ChannelID: channelID, UserID: userID},
})
}
@@ -393,7 +393,7 @@ func buildVoiceLeave(channelID, userID int64) []byte {
// buildChannelCreate constructs a channel_create broadcast.
func buildChannelCreate(ch *db.Channel) []byte {
return buildJSON(wsMsg{
Type: "channel_create",
Type: MsgTypeChannelCreate,
Payload: channelPayload{
ID: ch.ID,
Name: ch.Name,
@@ -408,7 +408,7 @@ func buildChannelCreate(ch *db.Channel) []byte {
// buildChannelUpdate constructs a channel_update broadcast.
func buildChannelUpdate(ch *db.Channel) []byte {
return buildJSON(wsMsg{
Type: "channel_update",
Type: MsgTypeChannelUpdate,
Payload: channelPayload{
ID: ch.ID,
Name: ch.Name,
@@ -423,7 +423,7 @@ func buildChannelUpdate(ch *db.Channel) []byte {
// buildChannelDelete constructs a channel_delete broadcast.
func buildChannelDelete(channelID int64) []byte {
return buildJSON(wsMsg{
Type: "channel_delete",
Type: MsgTypeChannelDelete,
Payload: channelDeletePayload{ID: channelID},
})
}
@@ -435,7 +435,7 @@ func buildDMChannelOpen(channelID int64, recipient *db.User) []byte {
avatarStr = *recipient.Avatar
}
return buildJSON(wsMsg{
Type: "dm_channel_open",
Type: MsgTypeDMChannelOpen,
Payload: dmChannelOpenPayload{
ChannelID: channelID,
Recipient: dmUserPayload{
@@ -451,7 +451,7 @@ func buildDMChannelOpen(channelID int64, recipient *db.User) []byte {
// buildServerRestartMsg constructs a server_restart broadcast.
func buildServerRestartMsg(reason string, delaySeconds int) []byte {
return buildJSON(wsMsg{
Type: "server_restart",
Type: MsgTypeServerRestart,
Payload: serverRestartPayload{
Reason: reason,
DelaySeconds: delaySeconds,
+48
View File
@@ -0,0 +1,48 @@
package ws
import "encoding/json"
// MessageHandler is the function signature for all WebSocket message handlers.
// It receives the hub, the sending client, the request ID from the envelope,
// and the raw JSON payload.
type MessageHandler func(h *Hub, c *Client, reqID string, payload json.RawMessage)
// HandlerRegistry maps message type strings to their handler functions.
// It is not safe for concurrent use after initialization; all Register
// calls must happen before any Dispatch calls.
type HandlerRegistry struct {
handlers map[string]MessageHandler
}
// NewHandlerRegistry creates an empty handler registry.
func NewHandlerRegistry() *HandlerRegistry {
return &HandlerRegistry{
handlers: make(map[string]MessageHandler),
}
}
// Register associates a message type with a handler function.
func (r *HandlerRegistry) Register(msgType string, handler MessageHandler) {
r.handlers[msgType] = handler
}
// Dispatch looks up the handler for msgType and invokes it. Returns true if a
// handler was found and called, false if no handler is registered for the type.
func (r *HandlerRegistry) Dispatch(msgType string, h *Hub, c *Client, reqID string, payload json.RawMessage) bool {
handler, ok := r.handlers[msgType]
if !ok {
return false
}
handler(h, c, reqID, payload)
return true
}
// RegisteredTypes returns all registered message types (unordered).
// Intended for testing and diagnostics.
func (r *HandlerRegistry) RegisteredTypes() []string {
types := make([]string, 0, len(r.handlers))
for t := range r.handlers {
types = append(types, t)
}
return types
}
+79
View File
@@ -0,0 +1,79 @@
package ws
import (
"encoding/json"
"sort"
"testing"
)
func TestHandlerRegistry_RegisterAndDispatch(t *testing.T) {
r := NewHandlerRegistry()
called := false
r.Register("test_type", func(h *Hub, c *Client, reqID string, payload json.RawMessage) {
called = true
if reqID != "req-1" {
t.Errorf("expected reqID %q, got %q", "req-1", reqID)
}
})
ok := r.Dispatch("test_type", nil, nil, "req-1", nil)
if !ok {
t.Fatal("Dispatch returned false for registered type")
}
if !called {
t.Fatal("handler was not called")
}
}
func TestHandlerRegistry_DispatchUnknownType(t *testing.T) {
r := NewHandlerRegistry()
ok := r.Dispatch("nonexistent", nil, nil, "", nil)
if ok {
t.Fatal("Dispatch returned true for unregistered type")
}
}
func TestHandlerRegistry_AllExpectedTypesRegistered(t *testing.T) {
r := NewHandlerRegistry()
registerChatHandlers(r)
registerPresenceHandlers(r)
registerReactionHandlers(r)
registerVoiceHandlers(r)
registerPingHandler(r)
expected := []string{
"chat_send",
"chat_edit",
"chat_delete",
"reaction_add",
"reaction_remove",
"typing_start",
"presence_update",
"channel_focus",
"voice_join",
"voice_leave",
"voice_token_refresh",
"voice_mute",
"voice_deafen",
"voice_camera",
"voice_screenshare",
"ping",
}
registered := r.RegisteredTypes()
sort.Strings(registered)
sort.Strings(expected)
if len(registered) != len(expected) {
t.Fatalf("expected %d registered types, got %d\nexpected: %v\ngot: %v",
len(expected), len(registered), expected, registered)
}
for i, typ := range expected {
if registered[i] != typ {
t.Errorf("mismatch at index %d: expected %q, got %q", i, typ, registered[i])
}
}
}
+2 -2
View File
@@ -270,7 +270,7 @@ func (h *Hub) buildAuthOK(user *db.User, roleName string) []byte {
serverName, motd := h.getCachedSettings()
return buildJSON(map[string]any{
"type": "auth_ok",
"type": MsgTypeAuthOK,
"payload": map[string]any{
"user": map[string]any{
"id": user.ID,
@@ -350,7 +350,7 @@ func (h *Hub) buildReady(database *db.DB, userID int64) ([]byte, error) {
serverName, motd := h.getCachedSettings()
return buildJSON(map[string]any{
"type": "ready",
"type": MsgTypeReady,
"payload": map[string]any{
"channels": channelPayloads,
"members": members,