Compare commits

...
2013 changed files with 330826 additions and 93915 deletions
+2
View File
@@ -50,6 +50,8 @@
/app-dist-output/
/artifacts/
/desktop-shared-assets/
/desktop-modules/
/s3_payload/
/upload_staging/
@@ -154,6 +154,7 @@ jobs:
BUILD_VERSION=${{ needs.meta.outputs.build_version }}
SOURCE_SHA=${{ github.sha }}
SOURCE_DATE=${{ steps.source.outputs.date }}
FLUXER_SELF_HOSTED=true
APP_ASSETS_REF=ghcr.io/${{ env.GHCR_OWNER }}/fluxer-app-proxy-self-hosted:${{ needs.meta.outputs.build_version }}-assets
APP_ASSETS_PLATFORM=linux/amd64
cache-from: type=registry,ref=ghcr.io/${{ env.GHCR_OWNER }}/fluxer-app-proxy-self-hosted:buildcache-${{ matrix.platform }}
+1
View File
@@ -200,6 +200,7 @@ jobs:
tags: ghcr.io/${{ env.GHCR_OWNER }}/fluxer-app-proxy:${{ needs.meta.outputs.build_version }}-arm64
build-args: |
BUILD_VERSION=${{ needs.meta.outputs.build_version }}
PUBLIC_ASSET_BASE_URL=https://fluxerstatic.com
SOURCE_SHA=${{ github.sha }}
SOURCE_DATE=${{ steps.source.outputs.date }}
APP_ASSETS_REF=ghcr.io/${{ env.GHCR_OWNER }}/fluxer-app-proxy:${{ needs.meta.outputs.build_version }}-assets
+164
View File
@@ -103,11 +103,150 @@ jobs:
--step set_matrix
--skip-targets "${{ inputs.skip_targets }}"
shared_assets:
name: Build shared renderer assets
needs:
- meta
runs-on: ubuntu-24.04
environment: desktop-releases
timeout-minutes: 60
permissions:
contents: read
env:
CHANNEL: ${{ needs.meta.outputs.channel }}
BUILD_CHANNEL: ${{ needs.meta.outputs.build_channel }}
RELEASE_CHANNEL: ${{ needs.meta.outputs.build_channel }}
PUBLIC_RELEASE_CHANNEL: ${{ needs.meta.outputs.build_channel }}
VERSION: ${{ needs.meta.outputs.version }}
BUILD_VERSION: ${{ needs.meta.outputs.version }}
PUBLIC_BUILD_VERSION: ${{ needs.meta.outputs.version }}
SOURCE_SHA: ${{ needs.meta.outputs.source_sha }}
steps:
- name: Checkout CI helpers
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1
with:
ref: ${{ needs.meta.outputs.source_sha }}
path: _ci
- name: Checkout source
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1
with:
ref: ${{ needs.meta.outputs.source_sha }}
path: source
- name: Set up Rust toolchain (renderer wasm)
uses: dtolnay/rust-toolchain@02cb101ec7c40f2c49e1d9714d64511d8e1b74de
with:
toolchain: "1.98.1"
targets: wasm32-unknown-unknown
- name: Set workdir (Unix)
env:
SUBST_TARGET: ${{ github.workspace }}/source
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step set_workdir_unix
- name: Set up Node.js
uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020
with:
node-version: 26
- name: Set up pnpm
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step setup_pnpm
- name: Resolve pnpm store path (Unix)
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step resolve_pnpm_store_unix
- name: Cache pnpm store
uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9
with:
path: ${{ env.PNPM_STORE_PATH }}
key: ${{ runner.os }}-shared-renderer-pnpm-store-${{ hashFiles('source/pnpm-lock.yaml') }}
restore-keys: |
${{ runner.os }}-shared-renderer-pnpm-store-
- name: Cache cargo registry
uses: actions/cache@55cc8345863c7cc4c66a329aec7e433d2d1c52a9
with:
path: |
~/.cargo/registry
~/.cargo/git
key: ${{ runner.os }}-shared-renderer-cargo-registry-${{ hashFiles('source/Cargo.lock') }}
restore-keys: |
${{ runner.os }}-shared-renderer-cargo-registry-
- name: Install dependencies
working-directory: ${{ env.WORKDIR }}/fluxer_desktop
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step install_dependencies
- name: Update version
working-directory: ${{ env.WORKDIR }}/fluxer_desktop
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step update_version
- name: Set build channel
working-directory: ${{ env.WORKDIR }}/fluxer_desktop
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step set_build_channel
- name: Build shared renderer assets
working-directory: ${{ env.WORKDIR }}/fluxer_desktop
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step build_shared_assets
- name: Prepare shared renderer artifact
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step prepare_shared_assets
- name: Upload shared renderer artifact
uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02
with:
name: desktop-shared-${{ needs.meta.outputs.build_channel }}-${{ needs.meta.outputs.version }}-${{ needs.meta.outputs.source_sha }}
path: source/desktop-shared-assets
if-no-files-found: error
retention-days: 1
compression-level: 0
- name: Split renderer into desktop modules
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step split_modules
- name: Pack desktop modules
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step pack_modules
- name: Upload desktop module packages
uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02
with:
name: desktop-modules-${{ needs.meta.outputs.build_channel }}-${{ needs.meta.outputs.version }}-${{ needs.meta.outputs.source_sha }}
path: |
source/desktop-modules/classification.json
source/desktop-modules/*/module.json
source/desktop-modules/*/package.br
source/desktop-modules/*/package.br.sha256
if-no-files-found: error
retention-days: 1
compression-level: 0
build:
name: Build ${{ matrix.platform }} (${{ matrix.arch }})
needs:
- meta
- matrix
- shared_assets
runs-on: ${{ matrix.os }}
environment: desktop-releases
timeout-minutes: 180
@@ -130,6 +269,7 @@ jobs:
SOURCE_SHA: ${{ needs.meta.outputs.source_sha }}
DESKTOP_PLATFORM: ${{ matrix.platform }}
DESKTOP_ARCH: ${{ matrix.arch }}
FLUXER_MODULES: "1"
PLATFORM: ${{ matrix.platform }}
ARCH: ${{ matrix.arch }}
ELECTRON_ARCH: ${{ matrix.electron_arch }}
@@ -276,6 +416,17 @@ jobs:
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step set_build_channel
- name: Download shared renderer artifact
uses: actions/download-artifact@d3f86a106a0bac45b974a628896c90dbdf5c8093
with:
name: desktop-shared-${{ needs.meta.outputs.build_channel }}-${{ needs.meta.outputs.version }}-${{ needs.meta.outputs.source_sha }}
path: source/desktop-shared-assets
- name: Restore shared renderer assets
run: >-
cargo run --locked --quiet --manifest-path ${{ github.workspace }}/_ci/tools/ci/Cargo.toml -- build-desktop
--step restore_shared_assets
- name: Build Electron main process
working-directory: ${{ env.WORKDIR }}/fluxer_desktop
env:
@@ -505,11 +656,13 @@ jobs:
if: ${{ !cancelled() && needs.build.result == 'success' }}
needs:
- meta
- shared_assets
- build
runs-on: ubuntu-24.04-arm
environment: desktop-releases
timeout-minutes: 180
permissions:
actions: read
contents: read
env:
CHANNEL: ${{ needs.meta.outputs.build_channel }}
@@ -547,6 +700,17 @@ jobs:
cargo run --locked --quiet --manifest-path tools/ci/Cargo.toml -- build-desktop
--step build_payload
- name: Download desktop module packages
uses: actions/download-artifact@d3f86a106a0bac45b974a628896c90dbdf5c8093
with:
name: desktop-modules-${{ needs.meta.outputs.build_channel }}-${{ needs.meta.outputs.version }}-${{ needs.meta.outputs.source_sha }}
path: desktop-modules
- name: Build desktop module manifest
run: >-
cargo run --locked --quiet --manifest-path tools/ci/Cargo.toml -- build-desktop
--step build_module_manifest
- name: Prepare GitHub release assets
run: >-
cargo run --locked --quiet --manifest-path tools/ci/Cargo.toml -- build-desktop
+2
View File
@@ -44,6 +44,8 @@
/app-dist-output/
/artifacts/
/desktop-shared-assets/
/desktop-modules/
/s3_payload/
/upload_staging/
Generated
+33 -1
View File
@@ -1710,6 +1710,16 @@ version = "0.2.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d"
[[package]]
name = "filetime"
version = "0.2.29"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5c287a33c7f0a620c38e641e7f60827713987b3c0f26e8ddc9462cc69cf75759"
dependencies = [
"cfg-if",
"libc",
]
[[package]]
name = "find-msvc-tools"
version = "0.1.12"
@@ -1741,6 +1751,7 @@ dependencies = [
"aws-config",
"aws-sdk-s3",
"base64 0.23.1",
"brotli",
"bytes",
"chrono",
"clap",
@@ -1750,6 +1761,7 @@ dependencies = [
"serde",
"serde_json",
"sha2 0.11.0",
"tar",
"tempfile",
"tokio",
"walkdir",
@@ -2042,7 +2054,6 @@ dependencies = [
"hex",
"rand 0.10.2",
"reqwest",
"serde",
"serde_json",
"sha2 0.11.0",
"tokio",
@@ -4790,6 +4801,17 @@ version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417"
[[package]]
name = "tar"
version = "0.4.46"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3f6221d9a6003c78398e3b239969f352578258df48c8eb051caadae0015bc840"
dependencies = [
"filetime",
"libc",
"xattr",
]
[[package]]
name = "tempfile"
version = "3.27.0"
@@ -5818,6 +5840,16 @@ dependencies = [
"tls_codec",
]
[[package]]
name = "xattr"
version = "1.6.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "32e45ad4206f6d2479085147f02bc2ef834ac85886624a23575ae137c8aa8156"
dependencies = [
"libc",
"rustix",
]
[[package]]
name = "xmlparser"
version = "0.13.6"
+1
View File
@@ -671,6 +671,7 @@ services:
FLUXER_S3_SECRET_ACCESS_KEY: ${FLUXER_S3_SECRET_KEY:?set FLUXER_S3_SECRET_KEY in .env}
FLUXER_S3_BUCKET_UPLOADS: ${FLUXER_S3_BUCKET_UPLOADS:-}
FLUXER_STATIC_CDN_ENDPOINT: ${FLUXER_STATIC_CDN_ENDPOINT:-}
FLUXER_MEDIA_ENDPOINT: ${FLUXER_MEDIA_ENDPOINT:-}
DISCOVERY_UPSTREAM_URL: http://edge:8088/.well-known/fluxer
DISCOVERY_REFRESH_INTERVAL_MS: ${DISCOVERY_REFRESH_INTERVAL_MS:-}
PUBLIC_BOOTSTRAP_API_ENDPOINT: /api
File diff suppressed because it is too large Load Diff
+25
View File
@@ -0,0 +1,25 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
use crate::api::generated::snowflake;
use super::client::{AdminApiClient, ApiResult};
use super::types::ListGuildThreadsResponse;
impl AdminApiClient {
pub async fn list_guild_threads(&self, guild_id: &str) -> ApiResult<ListGuildThreadsResponse> {
let response = self
.generated()
.list_admin_guild_threads(&snowflake(guild_id))
.await
.map_err(|e| self.generated_error(e))?;
self.generated_value(response.into_inner())
}
pub async fn delete_thread_channel(&self, channel_id: &str) -> ApiResult<()> {
self.generated()
.delete_admin_thread_channel(&snowflake(channel_id))
.await
.map_err(|e| self.generated_error(e))?;
Ok(())
}
}
+1
View File
@@ -13,6 +13,7 @@ pub mod client;
pub mod codes;
pub mod discovery;
pub mod guild_assets;
pub mod guild_threads;
pub mod guilds;
pub mod instance_config;
pub mod jobs;
@@ -0,0 +1,35 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct GuildThreadMetadata {
pub archived: bool,
pub locked: bool,
pub auto_archive_duration: i32,
pub archive_timestamp: String,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct GuildThreadItem {
pub id: String,
#[serde(rename = "type")]
pub channel_type: i32,
#[serde(default)]
pub name: Option<String>,
#[serde(default)]
pub parent_id: Option<String>,
#[serde(default)]
pub owner_id: Option<String>,
#[serde(default)]
pub member_count: Option<i32>,
#[serde(default)]
pub message_count: Option<i32>,
#[serde(default)]
pub thread_metadata: Option<GuildThreadMetadata>,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct ListGuildThreadsResponse {
pub threads: Vec<GuildThreadItem>,
}
@@ -31,6 +31,8 @@ pub struct InstanceConfigResponse {
#[serde(default)]
pub captcha: CaptchaConfigResponse,
#[serde(default)]
pub channel_threads: ChannelThreadsConfigResponse,
#[serde(default)]
pub experiment_delivery: ExperimentDeliveryConfigResponse,
#[serde(default)]
pub billing: InstanceBillingResponse,
@@ -510,6 +512,8 @@ pub const DOMAIN_MIGRATION_DEFAULT_SALT: &str = "domain-migration-v1";
pub const PLUTONIUM_PAGE_DEFAULT_SALT: &str = "plutonium-page-v1";
pub const CAPTCHA_COST_RANGE: std::ops::RangeInclusive<u32> = 1_000..=20_000;
pub const CAPTCHA_MAX_COUNTER_RANGE: std::ops::RangeInclusive<u32> = 100..=20_000;
pub const CHANNEL_THREADS_DEFAULT_GUILD_SALT: &str = "channel-threads-guild-v1";
pub const CHANNEL_THREADS_DEFAULT_USER_SALT: &str = "channel-threads-user-v1";
#[derive(Clone, Debug, Default, Deserialize, Serialize)]
#[serde(default)]
@@ -653,6 +657,62 @@ pub struct CaptchaConfigUpdateRequest {
pub max_counter: Option<u32>,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(default)]
pub struct ChannelThreadsConfigResponse {
pub enabled: bool,
pub config_version: u64,
pub ever_enabled: bool,
pub guild_basis_points: u32,
pub guild_salt: String,
pub enabled_guild_ids: Vec<String>,
pub disabled_guild_ids: Vec<String>,
pub user_basis_points: u32,
pub user_salt: String,
pub included_user_ids: Vec<String>,
pub excluded_user_ids: Vec<String>,
}
impl Default for ChannelThreadsConfigResponse {
fn default() -> Self {
Self {
enabled: false,
config_version: 0,
ever_enabled: false,
guild_basis_points: 0,
guild_salt: CHANNEL_THREADS_DEFAULT_GUILD_SALT.to_owned(),
enabled_guild_ids: Vec::new(),
disabled_guild_ids: Vec::new(),
user_basis_points: 0,
user_salt: CHANNEL_THREADS_DEFAULT_USER_SALT.to_owned(),
included_user_ids: Vec::new(),
excluded_user_ids: Vec::new(),
}
}
}
#[derive(Clone, Debug, Default, Serialize)]
pub struct ChannelThreadsConfigUpdateRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub enabled: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub guild_basis_points: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub guild_salt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub enabled_guild_ids: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub disabled_guild_ids: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user_basis_points: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user_salt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub included_user_ids: Option<Vec<String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub excluded_user_ids: Option<Vec<String>>,
}
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(default)]
pub struct ExperimentDeliveryConfigResponse {
@@ -775,6 +835,8 @@ pub struct InstanceConfigUpdateRequest {
#[serde(skip_serializing_if = "Option::is_none")]
pub captcha: Option<CaptchaConfigUpdateRequest>,
#[serde(skip_serializing_if = "Option::is_none")]
pub channel_threads: Option<ChannelThreadsConfigUpdateRequest>,
#[serde(skip_serializing_if = "Option::is_none")]
pub experiment_delivery: Option<ExperimentDeliveryConfigUpdateRequest>,
#[serde(skip_serializing_if = "Option::is_none")]
pub billing: Option<InstanceBillingUpdateRequest>,
+2
View File
@@ -9,6 +9,7 @@ mod codes;
mod common;
mod discovery;
mod guild_assets;
mod guild_threads;
mod instance_billing;
mod instance_config;
mod jobs;
@@ -29,6 +30,7 @@ pub use codes::*;
pub use common::*;
pub use discovery::*;
pub use guild_assets::*;
pub use guild_threads::*;
pub use instance_billing::*;
pub use instance_config::*;
pub use jobs::*;
+27
View File
@@ -96,6 +96,33 @@ pub async fn render(
config, &guild, &stickers, csrf_token,
))
}
"threads" => {
if !acl::has_permission(admin_acls, acl::GUILD_LOOKUP) {
return None;
}
let threads = client
.list_guild_threads(guild_id)
.await
.map(|response| response.threads)
.map_err(|error| tracing::warn!(%error, guild_id, "admin API request failed: list guild threads"))
.unwrap_or_default();
let threads_enabled = client
.get_instance_config()
.await
.map(|instance| instance.channel_threads.enabled)
.map_err(
|error| tracing::warn!(%error, "admin API request failed: get instance config"),
)
.unwrap_or(false);
Some(tabs::threads::threads_tab(
config,
&guild,
&threads,
acl::has_permission(admin_acls, acl::MESSAGE_DELETE_ALL),
threads_enabled,
csrf_token,
))
}
"audit_log" | "audit-log" => {
if !acl::has_permission(admin_acls, acl::GUILD_AUDIT_LOG_VIEW) {
return None;
+10
View File
@@ -433,6 +433,16 @@ async fn dispatch_guild_action(
"Failed to delete sticker",
)
}
"delete_thread" => {
let Some(thread_id) = get("thread_id") else {
return FlashData::error("Thread ID is required");
};
action_result(
client.delete_thread_channel(&thread_id).await,
"Thread deleted",
"Failed to delete thread",
)
}
"trigger_archive" => {
let inc = form.bool_value("include_attachments");
action_result(
+65 -13
View File
@@ -7,19 +7,20 @@ use crate::{
AppBrandingConfigUpdateRequest, AppLegalConfigUpdateRequest,
AppPublicConfigUpdateRequest, AppRegistrationConfigUpdateRequest,
AppSetupConfigUpdateRequest, CAPTCHA_COST_RANGE, CAPTCHA_MAX_COUNTER_RANGE,
CaptchaConfigUpdateRequest, CreateRegistrationUrlRequest,
DomainMigrationConfigUpdateRequest, EXPERIMENT_MAX_TARGETED_USERS,
ExperimentDeliveryConfigUpdateRequest, GatewayRolloutConfigUpdateRequest,
GatewayRolloutMode, InstanceAttachmentDecayUpdateRequest,
InstanceBlueskyIntegrationUpdateRequest, InstanceBlueskyKeyIntegrationUpdateRequest,
InstanceConfigUpdateRequest, InstanceEmailIntegrationUpdateRequest,
InstanceEmailSmtpIntegrationUpdateRequest, InstanceEmailSmtpTestRequest,
InstanceGifIntegrationUpdateRequest, InstanceIntegrationsUpdateRequest,
InstanceMediaUpdateRequest, InstancePolicyUpdateRequest,
InstanceRegistrationConfigUpdateRequest, InstanceServicesUpdateRequest,
InstanceYoutubeIntegrationUpdateRequest, LimitConfigUpdateRequest, LimitRule,
LimitRuleFilters, PlutoniumPageConfigUpdateRequest, PremiumMode,
PushRelayConfigUpdateRequest, RegistrationMode, SsoConfigUpdateRequest, VoiceE2eeScope,
CaptchaConfigUpdateRequest, ChannelThreadsConfigUpdateRequest,
CreateRegistrationUrlRequest, DomainMigrationConfigUpdateRequest,
EXPERIMENT_MAX_TARGETED_USERS, ExperimentDeliveryConfigUpdateRequest,
GatewayRolloutConfigUpdateRequest, GatewayRolloutMode,
InstanceAttachmentDecayUpdateRequest, InstanceBlueskyIntegrationUpdateRequest,
InstanceBlueskyKeyIntegrationUpdateRequest, InstanceConfigUpdateRequest,
InstanceEmailIntegrationUpdateRequest, InstanceEmailSmtpIntegrationUpdateRequest,
InstanceEmailSmtpTestRequest, InstanceGifIntegrationUpdateRequest,
InstanceIntegrationsUpdateRequest, InstanceMediaUpdateRequest,
InstancePolicyUpdateRequest, InstanceRegistrationConfigUpdateRequest,
InstanceServicesUpdateRequest, InstanceYoutubeIntegrationUpdateRequest,
LimitConfigUpdateRequest, LimitRule, LimitRuleFilters,
PlutoniumPageConfigUpdateRequest, PremiumMode, PushRelayConfigUpdateRequest,
RegistrationMode, SsoConfigUpdateRequest, VoiceE2eeScope,
},
},
config::AdminConfig,
@@ -228,6 +229,11 @@ pub async fn instance_config_post(
Ok(update) => instance_config_result(client.update_instance_config(&update).await),
Err(message) => FlashData::error(message),
},
"update_channel_threads" => instance_config_result(
client
.update_instance_config(&build_channel_threads_update(&form))
.await,
),
"update_experiment_delivery" => match build_experiment_delivery_update(&form) {
Ok(update) => instance_config_result(client.update_instance_config(&update).await),
Err(message) => FlashData::error(message),
@@ -671,6 +677,28 @@ fn build_captcha_update(form: &MultiValueForm) -> Result<InstanceConfigUpdateReq
})
}
fn build_channel_threads_update(form: &MultiValueForm) -> InstanceConfigUpdateRequest {
let channel_threads = if form.bool_value("channel_threads_everyone") {
ChannelThreadsConfigUpdateRequest {
enabled: Some(true),
guild_basis_points: Some(EXPERIMENT_ROLLOUT_BASIS_POINTS_MAX),
user_basis_points: Some(EXPERIMENT_ROLLOUT_BASIS_POINTS_MAX),
disabled_guild_ids: Some(Vec::new()),
excluded_user_ids: Some(Vec::new()),
..Default::default()
}
} else {
ChannelThreadsConfigUpdateRequest {
enabled: Some(false),
..Default::default()
}
};
InstanceConfigUpdateRequest {
channel_threads: Some(channel_threads),
..Default::default()
}
}
fn build_experiment_delivery_update(
form: &MultiValueForm,
) -> Result<InstanceConfigUpdateRequest, String> {
@@ -1436,6 +1464,30 @@ mod tests {
}
}
#[test]
fn build_channel_threads_update_turns_threads_on_for_everyone() {
let form = MultiValueForm::parse(b"_csrf=token&channel_threads_everyone=true");
assert_eq!(
serde_json::to_value(build_channel_threads_update(&form)).expect("serializable update"),
serde_json::json!({"channel_threads": {
"enabled": true,
"guild_basis_points": 10000,
"user_basis_points": 10000,
"disabled_guild_ids": [],
"excluded_user_ids": [],
}})
);
}
#[test]
fn build_channel_threads_update_only_turns_threads_off_when_unchecked() {
let form = MultiValueForm::parse(b"_csrf=token");
assert_eq!(
serde_json::to_value(build_channel_threads_update(&form)).expect("serializable update"),
serde_json::json!({"channel_threads": {"enabled": false}})
);
}
#[test]
fn build_push_relay_update_reads_the_consent_checkbox() {
let unchecked = build_push_relay_update(&MultiValueForm::parse(b"_csrf=token"));
@@ -26,6 +26,7 @@ pub const GUILD_TABS: &[(&str, &str)] = &[
("archives", "Archives"),
("emojis", "Emojis"),
("stickers", "Stickers"),
("threads", "Threads"),
("audit_logs", "Admin Audit Logs"),
("audit_log", "Guild Audit Log"),
("reports", "Reports"),
@@ -208,6 +209,7 @@ fn guild_tab_visible(_config: &AdminConfig, tab_id: &str, admin_acls: &[String])
"overview" | "members" | "settings" | "features" | "moderation" => true,
"reports" => acl::has_permission(admin_acls, acl::REPORT_VIEW),
"emojis" | "stickers" => acl::has_permission(admin_acls, acl::ASSET_PURGE),
"threads" => acl::has_permission(admin_acls, acl::GUILD_LOOKUP),
"audit_logs" => acl::has_permission(admin_acls, acl::AUDIT_LOG_VIEW),
"audit_log" => acl::has_permission(admin_acls, acl::GUILD_AUDIT_LOG_VIEW),
"archives" => acl::has_any_permission(
@@ -11,6 +11,7 @@ pub mod overview;
pub mod reports;
pub mod settings;
pub mod stickers;
pub mod threads;
use crate::{api::types::GuildDetailInfo, utils::user_tag::user_tag};
@@ -33,7 +33,11 @@ fn channel_type_label(channel_type: i32) -> &'static str {
2 => "Voice",
4 => "Category",
5 => "Announcement",
13 => "Link",
11 => "Public thread",
12 => "Private thread",
15 => "Forum",
16 => "Media",
998 => "Link",
_ => "Unknown",
}
}
@@ -179,7 +183,7 @@ pub fn overview_tab(config: &AdminConfig, guild: &GuildDetailInfo, csrf_token: &
} @else {
div class="flex flex-col gap-2" {
@for channel in &sorted_channels {
@let is_link = channel.channel_type == 13;
@let is_link = channel.channel_type == 998;
@let parent = channel.parent_id.as_deref()
.and_then(|pid| channels_by_id.get(pid));
@let parent_nsfw_override = parent
@@ -0,0 +1,106 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
use crate::{
api::types::{GuildInfo, GuildThreadItem},
config::AdminConfig,
templates::components::{
badge::{BadgeVariant, badge},
form::{csrf_input, danger_button, form_actions, submit_button},
page_container::card_with_header,
table::{data_table, table_cell, table_row},
},
};
use maud::{Markup, html};
pub fn threads_tab(
config: &AdminConfig,
guild: &GuildInfo,
threads: &[GuildThreadItem],
can_delete: bool,
can_reindex: bool,
csrf_token: &str,
) -> Markup {
let base = &config.base_path;
html! {
@if can_reindex {
(card_with_header("Thread search index", html! {
form method="post"
action={(base) "/guilds/" (guild.id) "?tab=threads&action=refresh_search_index"}
class="w-full" {
(csrf_input(csrf_token))
input type="hidden" name="index_type" value="threads";
(form_actions(html! {
(submit_button("Refresh threads"))
}))
}
}))
}
(card_with_header(
&format!("Threads ({})", threads.len()),
html! {
@if threads.is_empty() {
p class="text-sm text-neutral-500" { "No threads found for this guild." }
} @else {
(data_table(
&["Thread", "Parent", "State", "Members", "Messages", ""],
html! {
@for thread in threads {
(thread_row(base, &guild.id, thread, can_delete, csrf_token))
}
},
))
}
},
))
}
}
fn thread_kind(channel_type: i32) -> &'static str {
match channel_type {
10 => "Announcement",
12 => "Private",
_ => "Public",
}
}
fn thread_state(thread: &GuildThreadItem) -> Markup {
let metadata = thread.thread_metadata.as_ref();
let archived = metadata.is_some_and(|metadata| metadata.archived);
let locked = metadata.is_some_and(|metadata| metadata.locked);
html! {
div class="flex flex-wrap gap-1" {
(badge(thread_kind(thread.channel_type), BadgeVariant::Default))
@if archived { (badge("Archived", BadgeVariant::Default)) }
@if locked { (badge("Locked", BadgeVariant::Default)) }
}
}
}
fn thread_row(
base: &str,
guild_id: &str,
thread: &GuildThreadItem,
can_delete: bool,
csrf_token: &str,
) -> Markup {
table_row(html! {
(table_cell(false, html! {
div class="font-medium" { (thread.name.as_deref().unwrap_or("")) }
div class="text-xs text-neutral-500" { "ID: " (thread.id) }
}))
(table_cell(true, html! { (thread.parent_id.as_deref().unwrap_or("")) }))
(table_cell(false, thread_state(thread)))
(table_cell(true, html! { (thread.member_count.unwrap_or(0)) }))
(table_cell(true, html! { (thread.message_count.unwrap_or(0)) }))
(table_cell(false, html! {
@if can_delete {
form method="post"
action={(base) "/guilds/" (guild_id) "?tab=threads&action=delete_thread"} {
(csrf_input(csrf_token))
input type="hidden" name="thread_id" value=(thread.id);
(danger_button("Delete thread"))
}
}
}))
})
}
@@ -4,7 +4,7 @@ use crate::{
api::types::{
AccountIdentityConfigResponse, AccountIdentityMode, AppPublicConfigResponse,
CAPTCHA_COST_RANGE, CAPTCHA_MAX_COUNTER_RANGE, CaptchaConfigResponse,
DOMAIN_MIGRATION_DEFAULT_SALT, DomainMigrationConfigResponse,
ChannelThreadsConfigResponse, DOMAIN_MIGRATION_DEFAULT_SALT, DomainMigrationConfigResponse,
EXPERIMENT_MAX_TARGETED_USERS, ExperimentDeliveryConfigResponse,
GatewayRolloutConfigResponse, InstanceConfigResponse, InstanceIntegrationsResponse,
InstanceMediaResponse, InstancePolicyResponse, InstanceRegistrationResponse,
@@ -190,8 +190,13 @@ pub fn instance_config_page(
"Gateway rollout behavior and the limit rules applied to users and guilds.",
html! {
(gateway_rollout_section(base, csrf_token, &instance_config.gateway_rollout))
(domain_migration_section(base, csrf_token, &instance_config.domain_migration))
(plutonium_page_section(base, csrf_token, &instance_config.plutonium_page))
@if !instance_config.self_hosted {
(domain_migration_section(base, csrf_token, &instance_config.domain_migration))
}
(channel_threads_section(base, csrf_token, &instance_config.channel_threads))
@if !instance_config.self_hosted {
(plutonium_page_section(base, csrf_token, &instance_config.plutonium_page))
}
(experiment_delivery_section(base, csrf_token, &instance_config.experiment_delivery))
@if let Some(limit_config) = limit_config {
(limit_config_section(base, limit_config))
@@ -1484,6 +1489,48 @@ fn captcha_section(base: &str, csrf_token: &str, captcha: &CaptchaConfigResponse
)
}
fn channel_threads_available_to_everyone(channel_threads: &ChannelThreadsConfigResponse) -> bool {
channel_threads.enabled
&& channel_threads.guild_basis_points >= 10_000
&& channel_threads.user_basis_points >= 10_000
&& channel_threads.disabled_guild_ids.is_empty()
&& channel_threads.excluded_user_ids.is_empty()
}
fn channel_threads_section(
base: &str,
csrf_token: &str,
channel_threads: &ChannelThreadsConfigResponse,
) -> Markup {
let everyone = channel_threads_available_to_everyone(channel_threads);
section_card_with_description(
"Threads and forums",
"Threads, forum channels and media channels in every community.",
html! {
form method="post" action={(base) "/instance-config?action=update_channel_threads"} {
(csrf_input(csrf_token))
div class="space-y-6" {
(checkbox(
"channel_threads_everyone",
"true",
"Available to everyone",
everyone,
true,
))
@if channel_threads.enabled && !everyone {
p class="text-xs text-neutral-500" {
"Currently on for part of this instance. Saving applies the setting above to everyone."
}
}
(form_actions(html! {
(submit_button("Save"))
}))
}
}
},
)
}
fn experiment_delivery_section(
base: &str,
csrf_token: &str,
@@ -2227,6 +2274,38 @@ mod tests {
assert!(!markup.contains("at the cap"));
}
#[test]
fn channel_threads_section_is_a_single_everyone_toggle() {
let everyone = ChannelThreadsConfigResponse {
enabled: true,
guild_basis_points: 10_000,
user_basis_points: 10_000,
..ChannelThreadsConfigResponse::default()
};
let on = channel_threads_section("/admin", "csrf", &everyone).into_string();
assert!(on.contains("action=update_channel_threads"));
assert!(on.contains("name=\"channel_threads_everyone\""));
assert!(on.contains("checked"));
assert!(!on.contains("basis_points"));
assert!(!on.contains("part of this instance"));
let partial = ChannelThreadsConfigResponse {
enabled: true,
enabled_guild_ids: vec!["1600000000000000001".to_owned()],
user_basis_points: 10_000,
..ChannelThreadsConfigResponse::default()
};
let partial = channel_threads_section("/admin", "csrf", &partial).into_string();
assert!(!partial.contains("checked"));
assert!(partial.contains("part of this instance"));
let off =
channel_threads_section("/admin", "csrf", &ChannelThreadsConfigResponse::default())
.into_string();
assert!(!off.contains("checked"));
assert!(!off.contains("part of this instance"));
}
#[test]
fn plutonium_page_section_shows_the_rollout_and_list_counts() {
let plutonium_page = PlutoniumPageConfigResponse {
+118
View File
@@ -433,6 +433,19 @@ fn deserialize_instance_config_response_with_unknown_keys() {
"max_counter": 1000,
"future_captcha_knob": 1
},
"channel_threads": {
"enabled": true,
"config_version": 3,
"ever_enabled": true,
"guild_basis_points": 0,
"guild_salt": "channel-threads-guild-v1",
"enabled_guild_ids": ["1600000000000000001"],
"disabled_guild_ids": [],
"user_basis_points": 10000,
"user_salt": "channel-threads-user-v1",
"included_user_ids": [],
"excluded_user_ids": []
},
"experiment_delivery": {"poll_interval_seconds": 300, "poll_jitter_percent": 15},
"registration": {
"mode": "open",
@@ -595,6 +608,9 @@ fn deserialize_instance_config_response_with_unknown_keys() {
assert!(resp.push_relay.relay_consent_accepted);
assert!(resp.captcha.enabled);
assert_eq!(resp.captcha.max_counter, 1000);
assert!(resp.channel_threads.enabled);
assert_eq!(resp.channel_threads.config_version, 3);
assert_eq!(resp.channel_threads.enabled_guild_ids.len(), 1);
assert_eq!(resp.experiment_delivery.poll_interval_seconds, 300);
assert!(resp.policy.single_community_guild_id.is_none());
assert_eq!(resp.policy.services.gif_enabled, Some(true));
@@ -662,6 +678,63 @@ fn deserialize_instance_config_response_with_unknown_keys() {
);
}
#[test]
fn deserialize_channel_threads_config() {
let config: types::ChannelThreadsConfigResponse = serde_json::from_str(
r#"{
"enabled": true,
"config_version": 12,
"ever_enabled": true,
"guild_basis_points": 50,
"guild_salt": "channel-threads-guild-v2",
"enabled_guild_ids": ["1600000000000000001"],
"disabled_guild_ids": ["1600000000000000002", "1600000000000000003"],
"user_basis_points": 10000,
"user_salt": "channel-threads-user-v1",
"included_user_ids": [],
"excluded_user_ids": ["1500000000000000001"],
"future_threads_knob": 1
}"#,
)
.expect("a channel threads config must deserialize");
assert!(config.enabled);
assert!(config.ever_enabled);
assert_eq!(config.config_version, 12);
assert_eq!(config.guild_basis_points, 50);
assert_eq!(config.guild_salt, "channel-threads-guild-v2");
assert_eq!(config.enabled_guild_ids, vec!["1600000000000000001"]);
assert_eq!(config.disabled_guild_ids.len(), 2);
assert_eq!(config.user_basis_points, 10000);
assert_eq!(config.excluded_user_ids, vec!["1500000000000000001"]);
let absent: types::ChannelThreadsConfigResponse =
serde_json::from_str("{}").expect("an api without the experiment still deserializes");
assert!(!absent.enabled);
assert!(!absent.ever_enabled);
assert_eq!(absent.guild_salt, types::CHANNEL_THREADS_DEFAULT_GUILD_SALT);
assert_eq!(absent.user_salt, types::CHANNEL_THREADS_DEFAULT_USER_SALT);
}
#[test]
fn serialize_channel_threads_update_never_sends_server_owned_fields() {
let update = types::InstanceConfigUpdateRequest {
channel_threads: Some(types::ChannelThreadsConfigUpdateRequest {
enabled: Some(true),
enabled_guild_ids: Some(vec!["1600000000000000001".to_owned()]),
..Default::default()
}),
..Default::default()
};
assert_eq!(
serde_json::to_value(&update).unwrap(),
serde_json::json!({"channel_threads": {
"enabled": true,
"enabled_guild_ids": ["1600000000000000001"],
}})
);
}
#[test]
fn deserialize_push_relay_config() {
let accepted: types::PushRelayConfigResponse = serde_json::from_str(
@@ -1084,3 +1157,48 @@ fn account_identity_lock_is_unknown_when_the_api_omits_it() {
assert_eq!(identity.mode, types::AccountIdentityMode::Email);
assert_eq!(identity.locked, None);
}
#[test]
fn deserialize_guild_threads_response() {
let json = r#"{
"threads": [
{
"id": "1600000000000000010",
"type": 12,
"guild_id": "1600000000000000001",
"parent_id": "1600000000000000002",
"owner_id": "1500000000000000001",
"name": "secret plans",
"last_message_id": null,
"last_pin_timestamp": null,
"rate_limit_per_user": 0,
"flags": 0,
"thread_metadata": {
"archived": true,
"auto_archive_duration": 4320,
"archive_timestamp": "2026-09-27T12:00:00.000Z",
"locked": false,
"invitable": false,
"create_timestamp": "2026-09-26T12:00:00.000Z"
},
"message_count": 3,
"total_message_sent": 4,
"member_count": 2
}
]
}"#;
let generated: generated_types::ListGuildThreadsResponse =
serde_json::from_str(json).expect("the generated client must accept the thread list");
assert_eq!(generated.threads.len(), 1);
let resp: types::ListGuildThreadsResponse = serde_json::from_str(json).unwrap();
let thread = &resp.threads[0];
assert_eq!(thread.channel_type, 12);
assert_eq!(thread.name.as_deref(), Some("secret plans"));
assert_eq!(thread.member_count, Some(2));
assert!(
thread
.thread_metadata
.as_ref()
.is_some_and(|m| m.archived && !m.locked)
);
}
+58
View File
@@ -469,6 +469,7 @@ async fn mutating_admin_pages_render_usable_csrf_tokens() {
"/instance-config?action=update_gateway_rollout",
"/instance-config?action=update_sso",
"/instance-config?action=update_domain_migration",
"/instance-config?action=update_channel_threads",
"/instance-config?action=update_plutonium_page",
"/instance-config?action=update_experiment_delivery",
][..],
@@ -485,6 +486,50 @@ async fn mutating_admin_pages_render_usable_csrf_tokens() {
}
}
#[tokio::test]
async fn channel_threads_section_renders_and_saves_through_htmx_toasts() {
let app = setup().await;
let (headers, body) = get_with_headers(&app, "/instance-config", &[]).await;
assert_full_layout(&body);
assert!(body.contains("Threads and forums"), "{body}");
assert!(body.contains("Available to everyone"), "{body}");
assert!(
body.contains(r#"name="channel_threads_everyone""#),
"{body}"
);
assert!(
!body.contains("channel_threads_guild_basis_points"),
"{body}"
);
let csrf_token = csrf_cookie(&headers)
.unwrap_or_else(|| panic!("instance config page did not set csrf_token cookie\n{body}"));
let cookie = format!("{}; csrf_token={}", app.session_cookie, csrf_token);
let htmx_headers = [
("HX-Request", "true"),
("HX-Target", "flash-container"),
("Cookie", cookie.as_str()),
];
for form in [
format!("_csrf={csrf_token}&channel_threads_everyone=true"),
format!("_csrf={csrf_token}"),
] {
let (status, response_headers, response_body) = post_form_with_headers(
&app,
"/instance-config?action=update_channel_threads",
&htmx_headers,
&form,
)
.await;
assert_eq!(status, StatusCode::NO_CONTENT, "{response_body}");
let toast = response_headers
.get("X-Fluxer-Admin-Toast")
.and_then(|value| value.to_str().ok())
.unwrap_or_else(|| panic!("missing toast header\n{response_body}"));
assert!(toast.contains("Instance config updated"), "{toast}");
}
}
#[tokio::test]
async fn instance_config_registration_tables_show_copyable_urls_and_compact_pending_actions() {
let app = setup().await;
@@ -1193,6 +1238,19 @@ fn instance_config() -> Value {
"anonymous_rollout_basis_points": 0,
"standalone_forwarding": false
},
"channel_threads": {
"enabled": false,
"config_version": 3,
"ever_enabled": true,
"guild_basis_points": 0,
"guild_salt": "channel-threads-guild-v1",
"enabled_guild_ids": ["1600000000000000001"],
"disabled_guild_ids": [],
"user_basis_points": 10000,
"user_salt": "channel-threads-user-v1",
"included_user_ids": [],
"excluded_user_ids": ["1500000000000000009"]
},
"plutonium_page": {
"enabled": false,
"config_version": 0,
@@ -1,7 +1,14 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
type ElasticsearchFieldType = 'text' | 'keyword' | 'boolean' | 'long' | 'integer' | 'date' | 'float';
export type FluxerSearchIndexName = 'messages' | 'guilds' | 'users' | 'reports' | 'audit_logs' | 'guild_members';
export type FluxerSearchIndexName =
| 'messages'
| 'guilds'
| 'users'
| 'reports'
| 'audit_logs'
| 'guild_members'
| 'threads';
export interface ElasticsearchFieldMapping {
type: ElasticsearchFieldType;
@@ -173,6 +180,26 @@ export const ELASTICSEARCH_INDEX_DEFINITIONS: Record<FluxerSearchIndexName, Elas
},
},
},
threads: {
indexName: 'threads',
mappings: {
properties: {
id: keyword(),
guildId: keyword(),
parentId: keyword(),
type: integer(),
name: textWithKeyword(),
ownerId: keyword(),
archived: bool(),
locked: bool(),
appliedTagIds: keyword(),
createdAt: long(),
idSequence: long(),
lastMessageAt: long(),
archivedAt: long(),
},
},
},
audit_logs: {
indexName: 'audit_logs',
mappings: {
@@ -0,0 +1,97 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {Client} from '@elastic/elasticsearch';
import type {SortCombinations} from '@elastic/elasticsearch/lib/api/types';
import type {
SearchableThread,
ThreadSearchCursor,
ThreadSearchFilters,
} from '@fluxer/schema/src/contracts/search/SearchDocumentTypes';
import type {ElasticsearchDistributedLock} from '@pkgs/elasticsearch_search/src/adapters/ElasticsearchIndexAdapter';
import {ElasticsearchIndexAdapter} from '@pkgs/elasticsearch_search/src/adapters/ElasticsearchIndexAdapter';
import type {ElasticsearchFilter} from '@pkgs/elasticsearch_search/src/ElasticsearchFilterUtils';
import {
compactFilters,
esAndTerms,
esRangeFilter,
esTermFilter,
esTermsFilter,
} from '@pkgs/elasticsearch_search/src/ElasticsearchFilterUtils';
import {ELASTICSEARCH_INDEX_DEFINITIONS} from '@pkgs/elasticsearch_search/src/ElasticsearchIndexDefinitions';
const SORT_FIELDS = {
last_message_time: 'lastMessageAt',
archive_time: 'archivedAt',
creation_time: 'createdAt',
} as const;
function cursorFilter(cursor: ThreadSearchCursor, op: 'gt' | 'lt'): ElasticsearchFilter {
return {
bool: {
should: [
esRangeFilter('createdAt', {[op]: cursor.createdAt}),
{
bool: {
filter: [
esTermFilter('createdAt', cursor.createdAt),
esRangeFilter('idSequence', {[op]: cursor.idSequence}),
],
},
},
],
minimum_should_match: 1,
},
};
}
function buildThreadFilters(filters: ThreadSearchFilters): Array<ElasticsearchFilter | undefined> {
const clauses: Array<ElasticsearchFilter | undefined> = [
esTermFilter('guildId', filters.guildId),
esTermFilter('parentId', filters.parentId),
];
if (filters.publicOnly) {
clauses.push(
filters.privateThreadIds && filters.privateThreadIds.length > 0
? {
bool: {
should: [esTermsFilter('type', [10, 11]), esTermsFilter('id', filters.privateThreadIds)],
minimum_should_match: 1,
},
}
: esTermsFilter('type', [10, 11]),
);
}
if (filters.archived !== undefined) clauses.push(esTermFilter('archived', filters.archived));
if (filters.tagIds && filters.tagIds.length > 0) {
if (filters.tagSetting === 'match_all') clauses.push(...esAndTerms('appliedTagIds', filters.tagIds));
else clauses.push(esTermsFilter('appliedTagIds', filters.tagIds));
}
if (filters.after) clauses.push(cursorFilter(filters.after, 'gt'));
if (filters.before) clauses.push(cursorFilter(filters.before, 'lt'));
return compactFilters(clauses);
}
function buildThreadSort(filters: ThreadSearchFilters): Array<SortCombinations> | undefined {
const sortBy = filters.sortBy ?? 'last_message_time';
if (sortBy === 'relevance') return undefined;
const order = filters.sortOrder ?? 'desc';
return [...new Set([SORT_FIELDS[sortBy], 'createdAt', 'idSequence'])].map((field) => ({[field]: {order}}));
}
export interface ElasticsearchThreadAdapterOptions {
client: Client;
lock?: ElasticsearchDistributedLock;
}
export class ElasticsearchThreadAdapter extends ElasticsearchIndexAdapter<ThreadSearchFilters, SearchableThread> {
constructor(options: ElasticsearchThreadAdapterOptions) {
super({
client: options.client,
index: ELASTICSEARCH_INDEX_DEFINITIONS.threads,
searchableFields: ['name'],
buildFilters: buildThreadFilters,
buildSort: buildThreadSort,
lock: options.lock,
});
}
}
+2
View File
@@ -19,6 +19,7 @@ import {Hono} from 'hono';
interface CreateAPIAppOptions {
config: APIConfig;
logger: ILogger;
registerRoutes?: (routes: HonoApp) => void;
}
interface APIAppResult {
@@ -53,6 +54,7 @@ export async function createAPIApp(options: CreateAPIAppOptions): Promise<APIApp
routes.onError(TelemetryAwareAppErrorHandler);
routes.notFound(AppNotFoundHandler);
registerControllers(routes, config);
options.registerRoutes?.(routes);
const app = new Hono<HonoEnv>({strict: true});
const {middleware: metricsMiddleware, metricsHandler} = createMetricsMiddleware('api');
app.use('*', metricsMiddleware);
+1
View File
@@ -495,6 +495,7 @@ export function buildAPIConfigFromMaster(master: MasterConfig): APIConfig {
laneName: apiWorkerConfig?.lane,
taskName: apiWorkerConfig?.task as WorkerTaskName | undefined,
enableCronScheduler: apiWorkerConfig?.enable_cron_scheduler,
metricsPort: apiWorkerConfig?.metrics_port,
laneConcurrencyOverrides: {
realtime: apiWorkerConfig?.lane_concurrency_overrides?.realtime,
unfurl: apiWorkerConfig?.lane_concurrency_overrides?.unfurl,
+5
View File
@@ -12,6 +12,7 @@ import type {IGuildSearchService} from '@app/api/search/IGuildSearchService';
import type {IMessageSearchService} from '@app/api/search/IMessageSearchService';
import type {IReportSearchService} from '@app/api/search/IReportSearchService';
import type {ISearchProvider} from '@app/api/search/ISearchProvider';
import type {IThreadSearchService} from '@app/api/search/IThreadSearchService';
import type {IUserSearchService} from '@app/api/search/IUserSearchService';
import {DEFAULT_SEARCH_CLIENT_TIMEOUT_MS} from '@fluxer/constants/src/Timeouts';
import type {ElasticsearchDistributedLock} from '@pkgs/elasticsearch_search/src/adapters/ElasticsearchIndexAdapter';
@@ -87,6 +88,10 @@ export function getGuildMemberSearchService(): IGuildMemberSearchService | null
return searchProvider?.getGuildMemberSearchService() ?? null;
}
export function getThreadSearchService(): IThreadSearchService | null {
return searchProvider?.getThreadSearchService() ?? null;
}
export async function initializeSearch(lock?: ElasticsearchDistributedLock): Promise<void> {
if (searchProvider) {
await shutdownSearch();
+104 -6
View File
@@ -141,6 +141,8 @@ import {
type InviteRow,
PRIVATE_CHANNEL_COLUMNS,
type PrivateChannelRow,
READ_STATE_COLUMNS,
type ReadStateRow,
WEBHOOK_COLUMNS,
WEBHOOKS_BY_SOURCE_CHANNEL_COLUMNS,
type WebhookRow,
@@ -268,6 +270,30 @@ import {
type StorePurchaseByUserRow,
type StorePurchaseRow,
} from '@app/api/database/types/StoreBillingTypes';
import {
ACTIVE_THREADS_BY_GUILD_COLUMNS,
type ActiveThreadsByGuildRow,
ARCHIVED_THREADS_BY_PARENT_COLUMNS,
type ArchivedThreadsByParentRow,
FORUM_PINNED_THREAD_COLUMNS,
type ForumPinnedThreadRow,
GUILD_THREAD_STATE_COLUMNS,
type GuildThreadStateRow,
THREAD_MEMBER_COLUMNS,
THREAD_MEMBERS_BY_USER_COLUMNS,
THREAD_ONLY_CHANNELS_BY_GUILD_COLUMNS,
THREAD_PARENT_CONFIG_COLUMNS,
THREAD_STATE_COLUMNS,
THREAD_STATS_COLUMNS,
THREADS_BY_PARENT_COLUMNS,
type ThreadMemberRow,
type ThreadMembersByUserRow,
type ThreadOnlyChannelsByGuildRow,
type ThreadParentConfigRow,
type ThreadStateRow,
type ThreadStatsRow,
type ThreadsByParentRow,
} from '@app/api/database/types/ThreadTypes';
import {
FAVORITE_MEME_COLUMNS,
type FavoriteMemeRow,
@@ -579,6 +605,84 @@ export const DmStates = defineTable<DmStateRow, 'hi_user_id' | 'lo_user_id' | 'c
columns: DM_STATE_COLUMNS,
primaryKey: ['hi_user_id', 'lo_user_id', 'channel_id'],
});
export const ThreadState = defineTable<ThreadStateRow, 'thread_id'>({
name: 'thread_state',
columns: THREAD_STATE_COLUMNS,
primaryKey: ['thread_id'],
partitionKey: ['thread_id'],
});
export const ThreadStats = defineTable<ThreadStatsRow, 'thread_id'>({
name: 'thread_stats',
columns: THREAD_STATS_COLUMNS,
primaryKey: ['thread_id'],
partitionKey: ['thread_id'],
});
export const ThreadsByParent = defineTable<ThreadsByParentRow, 'parent_id' | 'thread_id', 'parent_id'>({
name: 'threads_by_parent',
columns: THREADS_BY_PARENT_COLUMNS,
primaryKey: ['parent_id', 'thread_id'],
partitionKey: ['parent_id'],
});
export const ActiveThreadsByGuild = defineTable<ActiveThreadsByGuildRow, 'guild_id' | 'thread_id', 'guild_id'>({
name: 'active_threads_by_guild',
columns: ACTIVE_THREADS_BY_GUILD_COLUMNS,
primaryKey: ['guild_id', 'thread_id'],
partitionKey: ['guild_id'],
});
export const ArchivedThreadsByParent = defineTable<
ArchivedThreadsByParentRow,
'parent_id' | 'is_private' | 'archive_timestamp' | 'thread_id',
'parent_id' | 'is_private'
>({
name: 'archived_threads_by_parent',
columns: ARCHIVED_THREADS_BY_PARENT_COLUMNS,
primaryKey: ['parent_id', 'is_private', 'archive_timestamp', 'thread_id'],
partitionKey: ['parent_id', 'is_private'],
});
export const ThreadMembers = defineTable<ThreadMemberRow, 'thread_id' | 'user_id', 'thread_id'>({
name: 'thread_members',
columns: THREAD_MEMBER_COLUMNS,
primaryKey: ['thread_id', 'user_id'],
partitionKey: ['thread_id'],
});
export const ThreadMembersByUser = defineTable<
ThreadMembersByUserRow,
'user_id' | 'guild_id' | 'parent_id' | 'is_private' | 'thread_id',
'user_id'
>({
name: 'thread_members_by_user',
columns: THREAD_MEMBERS_BY_USER_COLUMNS,
primaryKey: ['user_id', 'guild_id', 'parent_id', 'is_private', 'thread_id'],
partitionKey: ['user_id'],
});
export const ThreadParentConfig = defineTable<ThreadParentConfigRow, 'guild_id' | 'channel_id', 'guild_id'>({
name: 'thread_parent_config',
columns: THREAD_PARENT_CONFIG_COLUMNS,
primaryKey: ['guild_id', 'channel_id'],
partitionKey: ['guild_id'],
});
export const ForumPinnedThread = defineTable<ForumPinnedThreadRow, 'parent_id'>({
name: 'forum_pinned_thread',
columns: FORUM_PINNED_THREAD_COLUMNS,
primaryKey: ['parent_id'],
partitionKey: ['parent_id'],
});
export const ThreadOnlyChannelsByGuild = defineTable<
ThreadOnlyChannelsByGuildRow,
'guild_id' | 'channel_id',
'guild_id'
>({
name: 'thread_only_channels_by_guild',
columns: THREAD_ONLY_CHANNELS_BY_GUILD_COLUMNS,
primaryKey: ['guild_id', 'channel_id'],
partitionKey: ['guild_id'],
});
export const GuildThreadState = defineTable<GuildThreadStateRow, 'guild_id'>({
name: 'guild_thread_state',
columns: GUILD_THREAD_STATE_COLUMNS,
primaryKey: ['guild_id'],
partitionKey: ['guild_id'],
});
interface PinnedDmRow {
user_id: bigint;
@@ -593,12 +697,6 @@ export const PinnedDms = defineTable<PinnedDmRow, 'user_id' | 'channel_id'>({
primaryKey: ['user_id', 'channel_id'],
});
interface ReadStateRow {
user_id: bigint;
channel_id: bigint;
}
const READ_STATE_COLUMNS = ['user_id', 'channel_id'] as const satisfies ReadonlyArray<keyof ReadStateRow>;
export const ReadStates = defineTable<ReadStateRow, 'user_id' | 'channel_id'>({
name: 'read_states',
columns: READ_STATE_COLUMNS,
@@ -38,6 +38,7 @@ export const AdminAuditReadActions = {
LIST_GUILD_MEMBERS: 'list_guild_members',
LIST_GUILD_MEMORY_STATS: 'list_guild_memory_stats',
LIST_GUILD_STICKERS: 'list_guild_stickers',
LIST_GUILD_THREADS: 'list_guild_threads',
LIST_USER_APPLICATIONS: 'list_user_applications',
LIST_USER_CHANGE_LOG: 'list_user_change_log',
LIST_USER_DM_CHANNELS: 'list_user_dm_channels',
@@ -4,6 +4,7 @@ import {AdminAuditReadActions} from '@app/api/admin/AdminAuditActions';
import {recordAdminRead, recordAdminWrite} from '@app/api/admin/AdminAuditRecorder';
import {createUserID} from '@app/api/BrandedTypes';
import {Config} from '@app/api/Config';
import {enqueueRebuildThreadAutoArchiveQueue, enqueueThreadSearchBackfill} from '@app/api/channel/threads/ThreadJobs';
import {
type InstancePolicyConfig,
REGISTRATION_PENDING_APPROVAL_TRAIT,
@@ -14,6 +15,7 @@ import {requireAdminACL} from '@app/api/middleware/AdminMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
import {
getChannelThreadsConfigPublisher,
getGatewayRolloutConfigPublisher,
getInstanceConfigRepository,
getPushRelayConfigPublisher,
@@ -21,6 +23,11 @@ import {
import {RateLimitConfigs} from '@app/api/RateLimitConfig';
import type {HonoApp, HonoEnv} from '@app/api/types/HonoEnv';
import {Validator} from '@app/api/Validator';
import {
enabledThreadGuildIds,
enqueueThreadPermissionSeeds,
newlyEnabledThreadGuildIds,
} from '@app/api/worker/tasks/SeedThreadPermissions';
import {AdminACLs} from '@fluxer/constants/src/AdminACLs';
import {InputValidationError} from '@fluxer/errors/src/domains/core/InputValidationError';
import {InstancePolicyTransitionNotAllowedError} from '@fluxer/errors/src/domains/core/InstancePolicyTransitionNotAllowedError';
@@ -35,6 +42,10 @@ import {
PendingRegistrationActionRequest,
RegistrationUrlIdParam,
} from '@fluxer/schema/src/domains/admin/AdminSchemas';
import {
applyChannelThreadsConfigUpdate,
type ChannelThreadsConfig,
} from '@fluxer/schema/src/domains/admin/ChannelThreadsSchemas';
import {DomainMigrationConfigSchema} from '@fluxer/schema/src/domains/admin/DomainMigrationSchemas';
import {GatewayRolloutConfigSchema} from '@fluxer/schema/src/domains/admin/GatewayRolloutSchemas';
import {PlutoniumPageConfigSchema} from '@fluxer/schema/src/domains/admin/PlutoniumPageSchemas';
@@ -69,6 +80,7 @@ async function buildInstanceConfigResponse(): Promise<InstanceConfigResponse> {
domainMigration,
plutoniumPage,
captcha,
channelThreads,
experimentDelivery,
registrationConfig,
registrationUrls,
@@ -80,6 +92,7 @@ async function buildInstanceConfigResponse(): Promise<InstanceConfigResponse> {
instanceConfigRepository.getDomainMigrationConfig(),
instanceConfigRepository.getPlutoniumPageConfig(),
instanceConfigRepository.getCaptchaConfig(),
instanceConfigRepository.getChannelThreadsConfig(),
instanceConfigRepository.getExperimentDeliveryConfig(),
instanceConfigRepository.getRegistrationConfig(),
instanceConfigRepository.getRegistrationUrlsForAdmin(),
@@ -118,6 +131,7 @@ async function buildInstanceConfigResponse(): Promise<InstanceConfigResponse> {
domain_migration: domainMigration,
plutonium_page: plutoniumPage,
captcha,
channel_threads: channelThreads,
experiment_delivery: experimentDelivery,
registration: {
...registrationConfig,
@@ -418,6 +432,25 @@ export function InstanceConfigAdminController(app: HonoApp) {
await instanceConfigRepository.updateCaptchaConfig(patch);
}
}
let channelThreadsConfigVersion: number | undefined;
if (data.channel_threads) {
const patch = omitUndefinedFields(data.channel_threads);
if (Object.keys(patch).length > 0) {
let previous: ChannelThreadsConfig | undefined;
const landed = await instanceConfigRepository.updateChannelThreadsConfig((current) => {
previous = current;
return applyChannelThreadsConfigUpdate(current, patch);
});
channelThreadsConfigVersion = landed.config_version;
await enqueueThreadPermissionSeeds(ctx.get('channelRepository').threads, landed);
const newlyEnabled = previous ? newlyEnabledThreadGuildIds(previous, landed) : enabledThreadGuildIds(landed);
for (const guildId of newlyEnabled) {
await enqueueRebuildThreadAutoArchiveQueue(guildId);
await enqueueThreadSearchBackfill(guildId);
}
await getChannelThreadsConfigPublisher().publish(landed);
}
}
if (data.experiment_delivery) {
const patch = data.experiment_delivery;
await instanceConfigRepository.updateExperimentDeliveryConfig((current) =>
@@ -613,6 +646,7 @@ export function InstanceConfigAdminController(app: HonoApp) {
action: 'update_instance_config',
metadata: {
sections: listSuppliedSections(data),
channel_threads_config_version: channelThreadsConfigVersion?.toString(),
granted_acls: grantedSetupCompleterAdmin ? AdminACLs.WILDCARD : undefined,
},
});
@@ -26,7 +26,7 @@ export function SearchAdminController(app: HonoApp) {
operationId: 'create_admin_search_index_refresh',
summary: 'Refresh a search index',
description:
'Trigger a full or partial rebuild of the named search index. Creates a background job and returns its refresh ID for status tracking. The channel_messages and guild_members indexes are rebuilt one guild at a time and require guild_id, and favorite_memes requires user_id. Requires GUILD_LOOKUP permission.',
'Trigger a full or partial rebuild of the named search index. Creates a background job and returns its refresh ID for status tracking. The channel_messages, guild_members and threads indexes are rebuilt one guild at a time and require guild_id, and favorite_memes requires user_id. Requires GUILD_LOOKUP permission.',
responseSchema: RefreshSearchIndexResponse,
statusCode: 200,
security: 'adminApiKey',
@@ -0,0 +1,84 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {AdminAuditReadActions} from '@app/api/admin/AdminAuditActions';
import {recordAdminRead, recordAdminWrite} from '@app/api/admin/AdminAuditRecorder';
import {createChannelID, createGuildID} from '@app/api/BrandedTypes';
import {requireAdminACL} from '@app/api/middleware/AdminMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
import {AdminRateLimitConfigs} from '@app/api/rate_limit_configs/AdminRateLimitConfig';
import type {HonoApp} from '@app/api/types/HonoEnv';
import {Validator} from '@app/api/Validator';
import {AdminACLs} from '@fluxer/constants/src/AdminACLs';
import {InvalidChannelTypeError} from '@fluxer/errors/src/domains/channel/InvalidChannelTypeError';
import {UnknownChannelError} from '@fluxer/errors/src/domains/channel/UnknownChannelError';
import {ListGuildThreadsResponse} from '@fluxer/schema/src/domains/admin/AdminThreadSchemas';
import {ChannelIdParam, GuildIdParam} from '@fluxer/schema/src/domains/common/CommonParamSchemas';
export function ThreadAdminController(app: HonoApp) {
app.get(
'/admin/guilds/:guild_id/threads',
RateLimitMiddleware(AdminRateLimitConfigs.ADMIN_LOOKUP),
requireAdminACL(AdminACLs.GUILD_LOOKUP),
Validator('param', GuildIdParam),
OpenAPI({
operationId: 'list_admin_guild_threads',
summary: 'List guild threads',
description:
'Lists every thread of a guild, active and archived, whether or not the channel threads experiment is active for it. Requires GUILD_LOOKUP permission.',
responseSchema: ListGuildThreadsResponse,
statusCode: 200,
security: 'adminApiKey',
tags: 'Admin',
experiment: 'channel_threads',
}),
async (ctx) => {
const guildId = createGuildID(ctx.req.valid('param').guild_id);
const threads = await ctx.get('threadService').lists.listGuildThreadsForAdmin(guildId);
await recordAdminRead(ctx, {
targetType: 'guild',
targetId: guildId,
action: AdminAuditReadActions.LIST_GUILD_THREADS,
metadata: {result_count: threads.length},
});
return ctx.json({threads});
},
);
app.delete(
'/admin/channels/:channel_id',
RateLimitMiddleware(AdminRateLimitConfigs.ADMIN_MESSAGE_OPERATION),
requireAdminACL(AdminACLs.MESSAGE_DELETE_ALL),
Validator('param', ChannelIdParam),
OpenAPI({
operationId: 'delete_admin_thread_channel',
summary: 'Delete a thread',
description:
'Deletes a thread channel with its messages and memberships. Only public and private threads can be deleted here. Requires MESSAGE_DELETE_ALL permission.',
responseSchema: null,
statusCode: 204,
security: 'adminApiKey',
tags: 'Admin',
experiment: 'channel_threads',
}),
async (ctx) => {
const channelId = createChannelID(ctx.req.valid('param').channel_id);
const thread = await ctx.get('channelRepository').findUnique(channelId);
if (!thread) throw new UnknownChannelError();
if (!thread.isThread()) throw new InvalidChannelTypeError();
const adminUserId = ctx.get('adminUserId');
await ctx.get('threadService').deletion.deleteThread({
thread,
actorId: adminUserId,
auditLogReason: ctx.get('auditLogReason'),
recordGuildAudit: false,
});
await recordAdminWrite(ctx, {
targetType: 'channel',
targetId: channelId,
action: 'delete_thread',
metadata: {guild_id: thread.guildId?.toString(), parent_id: thread.parentId?.toString(), type: thread.type},
});
return ctx.body(null, 204);
},
);
}
@@ -19,6 +19,7 @@ import {ReportAdminController} from '@app/api/admin/controllers/ReportAdminContr
import {SearchAdminController} from '@app/api/admin/controllers/SearchAdminController';
import {StoreBillingAdminController} from '@app/api/admin/controllers/StoreBillingAdminController';
import {SystemDmAdminController} from '@app/api/admin/controllers/SystemDmAdminController';
import {ThreadAdminController} from '@app/api/admin/controllers/ThreadAdminController';
import {UserAdminController} from '@app/api/admin/controllers/UserAdminController';
import {VoiceAdminController} from '@app/api/admin/controllers/VoiceAdminController';
import type {HonoApp} from '@app/api/types/HonoEnv';
@@ -30,6 +31,7 @@ export function registerAdminControllers(app: HonoApp) {
StoreBillingAdminController(app);
CodesAdminController(app);
GuildAdminController(app);
ThreadAdminController(app);
AssetAdminController(app);
BanAdminController(app);
InstanceConfigAdminController(app);
@@ -13,17 +13,19 @@ import {
type UserID,
} from '@app/api/BrandedTypes';
import type {IChannelRepository} from '@app/api/channel/IChannelRepository';
import {withThreadContext} from '@app/api/channel/services/ChannelGatewayDispatch';
import {
enqueueCrosspostFamilyPurgeFromCopies,
enqueueCrosspostSourceRemoval,
} from '@app/api/channel/services/message/CrosspostPropagation';
import {purgeMessageAttachments} from '@app/api/channel/services/message/MessageHelpers';
import {decrementThreadMessageCount, purgeMessageAttachments} from '@app/api/channel/services/message/MessageHelpers';
import {
createMessageResponseDataService,
type MessageResponseAccessContext,
messageResponseAccessForChannel,
messageResponseAccessForGuild,
} from '@app/api/channel/services/message/MessageResponseDataService';
import {resolveNsfwScopeChannel} from '@app/api/channel/utils/ThreadNsfwScope';
import type {NcmecAttachmentStatusResponse, NcmecSubmissionService} from '@app/api/csam/NcmecSubmissionService';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
import {getPurgeQueue, getStorageService} from '@app/api/middleware/ServiceSingletons';
@@ -144,15 +146,16 @@ export class AdminMessageService {
message.authorId || createUserID(0n),
message.pinnedTimestamp || undefined,
);
await decrementThreadMessageCount(channelRepository, channel, [messageId]);
if (channel) {
if (channel.guildId) {
await gatewayService.dispatchGuild({
guildId: channel.guildId,
event: 'MESSAGE_DELETE',
data: {
data: withThreadContext(channel, {
channel_id: channelId.toString(),
id: messageId.toString(),
},
}),
});
} else {
for (const recipientId of channel.recipientIds) {
@@ -349,9 +352,12 @@ export class AdminMessageService {
guildName: null,
};
}
const guild = await guildRepository.findUnique(channel.guildId);
const [guild, scope] = await Promise.all([
guildRepository.findUnique(channel.guildId),
resolveNsfwScopeChannel(channel, (id) => channelRepository.findUnique(id)),
]);
return {
channelNsfw: channel.isNsfw,
channelNsfw: scope.isNsfw,
guildNsfwLevel: guild?.nsfwLevel ?? null,
channelName: channel.name ?? null,
guildId: channel.guildId.toString(),
@@ -21,9 +21,11 @@ import {
messageResponseAccessForChannel,
messageResponseAccessForGuild,
} from '@app/api/channel/services/message/MessageResponseDataService';
import {resolveNsfwScopeChannel} from '@app/api/channel/utils/ThreadNsfwScope';
import {SYSTEM_USER_ID} from '@app/api/constants/Core';
import type {NcmecAttachmentStatusResponse, NcmecSubmissionService} from '@app/api/csam/NcmecSubmissionService';
import type {MessageAttachment} from '@app/api/database/types/MessageTypes';
import {SYSTEM_THREAD_VIEWER} from '@app/api/experiment/ChannelThreadsGate';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
import type {IStorageService} from '@app/api/infrastructure/IStorageService';
import type {UserCacheService} from '@app/api/infrastructure/UserCacheService';
@@ -205,6 +207,7 @@ export class AdminReportService {
});
await this.deps.channelService.messages.send.sendMessage({
user: systemUser,
viewer: SYSTEM_THREAD_VIEWER,
channelId: dmChannel.id,
data: {
content: template.value.body,
@@ -523,7 +526,10 @@ export class AdminReportService {
return reportNsfwLookupCache.channelNsfwByChannelId.get(channelIdString) ?? null;
}
const channel = await this.deps.channelRepository.findUnique(channelId);
const channelNsfw = channel?.isNsfw ?? null;
const scope = channel
? await resolveNsfwScopeChannel(channel, (id) => this.deps.channelRepository.findUnique(id))
: null;
const channelNsfw = scope?.isNsfw ?? null;
reportNsfwLookupCache.channelNsfwByChannelId.set(channelIdString, channelNsfw);
return channelNsfw;
}
@@ -6,9 +6,11 @@ import {mapUserToAdminResponse} from '@app/api/admin/models/UserTypes';
import type {AdminAuditService} from '@app/api/admin/services/AdminAuditService';
import {createGuildID, createUserID, type UserID} from '@app/api/BrandedTypes';
import {isSyntheticUserId} from '@app/api/constants/Core';
import {channelThreadsEnabled} from '@app/api/experiment/ChannelThreadsGate';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
import {Logger} from '@app/api/Logger';
import {getGuildSearchService, getUserSearchService} from '@app/api/SearchFactory';
import {ValidationErrorCodes} from '@fluxer/constants/src/ValidationErrorCodes';
import {FeatureTemporarilyDisabledError} from '@fluxer/errors/src/domains/core/FeatureTemporarilyDisabledError';
import {InputValidationError} from '@fluxer/errors/src/domains/core/InputValidationError';
import type {UserSearchFilters} from '@fluxer/schema/src/contracts/search/SearchDocumentTypes';
@@ -23,7 +25,8 @@ interface RefreshSearchIndexJobPayload extends WorkerJobPayload {
| 'channel_messages'
| 'favorite_memes'
| 'guild_members'
| 'discovery';
| 'discovery'
| 'threads';
admin_user_id: string;
audit_log_reason: string | null;
job_id: string;
@@ -175,7 +178,8 @@ export class AdminSearchService {
| 'channel_messages'
| 'guild_members'
| 'favorite_memes'
| 'discovery';
| 'discovery'
| 'threads';
guild_id?: bigint;
user_id?: bigint;
},
@@ -197,6 +201,15 @@ export class AdminSearchService {
}
payload.guild_id = data.guild_id.toString();
}
if (data.index_type === 'threads') {
if (!channelThreadsEnabled()) {
throw InputValidationError.fromCode('index_name', ValidationErrorCodes.INVALID_FORMAT);
}
if (!data.guild_id) {
throw InputValidationError.create('guild_id', 'guild_id is required for the threads index type');
}
payload.guild_id = data.guild_id.toString();
}
if (data.index_type === 'guild_members') {
if (!data.guild_id) {
throw InputValidationError.create('guild_id', 'guild_id is required for the guild_members index type');
@@ -39,7 +39,7 @@ export class AdminGuildLookupService {
return {guild: null};
}
const [channels, roles, ownerUser] = await Promise.all([
channelRepository.listGuildChannels(guildId),
channelRepository.listGuildChannels(guildId, 'maintenance'),
guildRepository.listRoles(guildId),
userRepository.findUnique(guild.ownerId),
]);
@@ -0,0 +1,83 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {createTestAccount, setUserACLs} from '@app/api/auth/tests/AuthTestUtils';
import {createChannel, createGuild} from '@app/api/channel/tests/ChannelTestUtils';
import {
ALL_THREADS_ACTIVE,
resetChannelThreadsConfig,
setChannelThreadsConfig,
threadsRequest,
} from '@app/api/channel/tests/ThreadTestUtils';
import {ensureSessionStarted, sendMessage} from '@app/api/message/tests/MessageTestUtils';
import {type ApiTestHarness, createApiTestHarness} from '@app/api/test/ApiTestHarness';
import {createBuilder} from '@app/api/test/TestRequestBuilder';
import {AdminACLs} from '@fluxer/constants/src/AdminACLs';
import {ChannelTypes, MessageTypes} from '@fluxer/constants/src/ChannelConstants';
import {ServerMessageFlags} from '@fluxer/constants/src/ThreadConstants';
import type {ThreadChannelResponse} from '@fluxer/schema/src/domains/channel/ThreadRequestSchemas';
import type {MessageResponse} from '@fluxer/schema/src/domains/message/MessageResponseSchemas';
import {afterAll, beforeAll, beforeEach, describe, expect, it} from 'vitest';
interface BrowseResponse {
messages: Array<{id: string; channel_id: string}>;
message_responses?: Array<MessageResponse>;
}
describe('admin thread browse', () => {
let harness: ApiTestHarness;
beforeAll(async () => {
harness = await createApiTestHarness();
});
beforeEach(async () => {
await harness.reset();
resetChannelThreadsConfig();
});
afterAll(async () => {
resetChannelThreadsConfig();
await harness.shutdown();
});
it('returns thread messages and thread artifacts unmasked', async () => {
await setChannelThreadsConfig(ALL_THREADS_ACTIVE);
const owner = await createTestAccount(harness);
await ensureSessionStarted(harness, owner.token);
const guild = await createGuild(harness, owner.token, 'browse');
const channel = await createChannel(harness, owner.token, guild.id, 'general');
const source = await sendMessage(harness, owner.token, channel.id, 'source');
await threadsRequest(harness, owner.token)
.post(`/channels/${channel.id}/messages/${source.id}/threads`)
.body({name: 'from source'})
.expect(201)
.execute();
const standalone = await threadsRequest<ThreadChannelResponse>(harness, owner.token)
.post(`/channels/${channel.id}/threads`)
.body({name: 'standalone', type: ChannelTypes.PUBLIC_THREAD})
.expect(201)
.execute();
const inThread = await threadsRequest<MessageResponse>(harness, owner.token)
.post(`/channels/${standalone.id}/messages`)
.body({content: 'inside'})
.expect(200)
.execute();
const admin = await setUserACLs(harness, await createTestAccount(harness), [
AdminACLs.AUTHENTICATE,
AdminACLs.MESSAGE_LOOKUP,
]);
const parent = await createBuilder<BrowseResponse>(harness, admin.token)
.get(`/admin/channels/${channel.id}/messages?limit=50`)
.expect(200)
.execute();
const responses = parent.message_responses ?? [];
const sourceResponse = responses.find((message) => message.id === source.id);
expect((sourceResponse?.flags ?? 0) & ServerMessageFlags.HAS_THREAD).toBe(ServerMessageFlags.HAS_THREAD);
expect(responses.some((message) => message.type === MessageTypes.THREAD_CREATED)).toBe(true);
const thread = await createBuilder<BrowseResponse>(harness, admin.token)
.get(`/admin/channels/${standalone.id}/messages?limit=50`)
.expect(200)
.execute();
expect(thread.messages.map((message) => message.id)).toContain(inThread.id);
});
});
@@ -21,6 +21,7 @@ import {ReportAdminAuditCases} from '@app/api/admin/tests/audit_coverage/ReportA
import {SearchAdminAuditCases} from '@app/api/admin/tests/audit_coverage/SearchAdminAuditCases';
import {StoreBillingAdminAuditCases} from '@app/api/admin/tests/audit_coverage/StoreBillingAdminAuditCases';
import {SystemDmAdminAuditCases} from '@app/api/admin/tests/audit_coverage/SystemDmAdminAuditCases';
import {ThreadAdminAuditCases} from '@app/api/admin/tests/audit_coverage/ThreadAdminAuditCases';
import {UserAdminAuditCases} from '@app/api/admin/tests/audit_coverage/UserAdminAuditCases';
import {UserWriteAdminAuditCases} from '@app/api/admin/tests/audit_coverage/UserWriteAdminAuditCases';
import {VoiceAdminAuditCases} from '@app/api/admin/tests/audit_coverage/VoiceAdminAuditCases';
@@ -48,6 +49,7 @@ const ALL_CASES = [
...SearchAdminAuditCases,
...StoreBillingAdminAuditCases,
...SystemDmAdminAuditCases,
...ThreadAdminAuditCases,
...UserAdminAuditCases,
...UserWriteAdminAuditCases,
...VoiceAdminAuditCases,
@@ -8,11 +8,13 @@ import {
type TestAccount,
} from '@app/api/auth/tests/AuthTestUtils';
import {getPngDataUrl} from '@app/api/emoji/tests/EmojiTestUtils';
import {ChannelThreadsConfigPublisher} from '@app/api/instance/ChannelThreadsConfigPublisher';
import {getInstanceConfigRepository} from '@app/api/middleware/ServiceSingletons';
import type {ApiTestHarness} from '@app/api/test/ApiTestHarness';
import {HTTP_STATUS, TEST_IDS} from '@app/api/test/TestConstants';
import {createBuilder} from '@app/api/test/TestRequestBuilder';
import type {CreateRegistrationUrlResponse} from '@fluxer/schema/src/domains/admin/AdminSchemas';
import {vi} from 'vitest';
async function createRegistrationUrl(harness: ApiTestHarness, admin: TestAccount): Promise<string> {
const created = await createBuilder<CreateRegistrationUrlResponse>(harness, admin.token)
@@ -80,6 +82,26 @@ export const InstanceConfigAdminAuditCases: ReadonlyArray<AdminAuditCoverageCase
};
},
},
{
method: 'PATCH',
route: '/admin/instance/config',
name: 'channel threads',
async prepare() {
vi.spyOn(ChannelThreadsConfigPublisher.prototype, 'publish').mockResolvedValue(undefined);
return {
request: {
path: '/admin/instance/config',
body: {channel_threads: {enabled: true, enabled_guild_ids: ['1']}},
},
expected: {
action: 'update_instance_config',
targetType: 'instance_config',
targetId: '0',
metadata: {sections: 'channel_threads', channel_threads_config_version: '1'},
},
};
},
},
{
method: 'POST',
route: '/admin/instance/config/branding-assets',
@@ -0,0 +1,56 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {AdminAuditCoverageCase} from '@app/api/admin/tests/audit_coverage/AdminAuditCoverage';
import {createChannel, createGuild} from '@app/api/channel/tests/ChannelTestUtils';
import {
ALL_THREADS_ACTIVE,
resetChannelThreadsConfig,
setChannelThreadsConfig,
threadsRequest,
} from '@app/api/channel/tests/ThreadTestUtils';
import {ChannelTypes} from '@fluxer/constants/src/ChannelConstants';
import type {ThreadChannelResponse} from '@fluxer/schema/src/domains/channel/ThreadRequestSchemas';
export const ThreadAdminAuditCases: ReadonlyArray<AdminAuditCoverageCase> = [
{
method: 'GET',
route: '/admin/guilds/:guild_id/threads',
async prepare({harness, admin}) {
resetChannelThreadsConfig();
const guild = await createGuild(harness, admin.token, 'Audit Thread List Guild');
return {
request: {path: `/admin/guilds/${guild.id}/threads`},
expected: {
action: 'list_guild_threads',
targetType: 'guild',
targetId: guild.id,
metadata: {result_count: '0'},
},
};
},
},
{
method: 'DELETE',
route: '/admin/channels/:channel_id',
async prepare({harness, admin}) {
resetChannelThreadsConfig();
await setChannelThreadsConfig(ALL_THREADS_ACTIVE);
const guild = await createGuild(harness, admin.token, 'Audit Thread Delete Guild');
const channel = await createChannel(harness, admin.token, guild.id, 'general');
const thread = await threadsRequest<ThreadChannelResponse>(harness, admin.token)
.post(`/channels/${channel.id}/threads`)
.body({name: 'doomed', type: ChannelTypes.PUBLIC_THREAD})
.expect(201)
.execute();
return {
request: {path: `/admin/channels/${thread.id}`, expectStatus: 204},
expected: {
action: 'delete_thread',
targetType: 'channel',
targetId: thread.id,
metadata: {guild_id: guild.id, parent_id: channel.id, type: String(ChannelTypes.PUBLIC_THREAD)},
},
};
},
},
];
@@ -0,0 +1,6 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {describeAdminAuditCoverage} from '@app/api/admin/tests/audit_coverage/AdminAuditCoverage';
import {ThreadAdminAuditCases} from '@app/api/admin/tests/audit_coverage/ThreadAdminAuditCases';
describeAdminAuditCoverage('ThreadAdminController', ThreadAdminAuditCases);
+11 -2
View File
@@ -3,6 +3,7 @@
import type {ILogger} from '@app/api/ILogger';
import {ActivityContextMiddleware} from '@app/api/infrastructure/activity/ActivityMeta';
import {AuditLogMiddleware} from '@app/api/middleware/AuditLogMiddleware';
import {ClientFeaturesMiddleware} from '@app/api/middleware/ClientFeaturesMiddleware';
import {ConcurrencyLimitMiddleware} from '@app/api/middleware/ConcurrencyLimitMiddleware';
import ContentFilterMiddleware from '@app/api/middleware/ContentFilterMiddleware';
import {GuildAvailabilityMiddleware} from '@app/api/middleware/GuildAvailabilityMiddleware';
@@ -57,10 +58,17 @@ export function configureMiddleware(routes: HonoApp, options: MiddlewarePipeline
allowedHeaders: [
HttpHeaders.CONTENT_TYPE,
HttpHeaders.AUTHORIZATION,
'X-Requested-With',
'Accept-Language',
HttpHeaders.X_REQUESTED_WITH,
HttpHeaders.ACCEPT_LANGUAGE,
HttpHeaders.X_REQUEST_ID,
HttpHeaders.IF_NONE_MATCH,
HttpHeaders.X_AUDIT_LOG_REASON,
HttpHeaders.X_CAPTCHA_ID,
HttpHeaders.X_CAPTCHA_TOKEN,
HttpHeaders.X_FLUXER_CLIENT_INSTALLATION_ID,
HttpHeaders.X_FLUXER_FEATURES,
HttpHeaders.X_FLUXER_PLATFORM,
HttpHeaders.X_FLUXER_SUDO_MODE_JWT,
],
exposedHeaders: [HttpHeaders.X_FLUXER_VERSION, HttpHeaders.ETAG],
},
@@ -78,6 +86,7 @@ export function configureMiddleware(routes: HonoApp, options: MiddlewarePipeline
);
routes.use(RequestErrorTelemetry);
routes.use(RequestCacheMiddleware);
routes.use(ClientFeaturesMiddleware);
if (nodeEnv === 'production') {
routes.use('*', async (ctx, next) => {
const host = ctx.req.header('host');
@@ -99,6 +99,17 @@ function serializeGuildTextChannel(channel: Channel, ctx: ContentWarningCtx): Ch
};
}
function serializeThreadOnlyChannel(channel: Channel, ctx: ContentWarningCtx): ChannelResponse {
return {
...serializeBaseChannelFields(channel),
...serializeMessageableFields(channel),
...serializePositionableGuildChannelFields(channel),
topic: channel.topic,
...serializeContentWarningFields(channel, ctx),
rate_limit_per_user: channel.rateLimitPerUser,
};
}
function serializeGuildVoiceChannel(channel: Channel, ctx: ContentWarningCtx): ChannelResponse {
return {
...serializeBaseChannelFields(channel),
@@ -202,6 +213,10 @@ export async function mapChannelToResponse(params: MapChannelToResponseParams):
case ChannelTypes.GUILD_VOICE:
response = serializeGuildVoiceChannel(channel, ctx);
break;
case ChannelTypes.GUILD_FORUM:
case ChannelTypes.GUILD_MEDIA:
response = serializeThreadOnlyChannel(channel, ctx);
break;
case ChannelTypes.GUILD_CATEGORY:
response = serializeGuildCategoryChannel(channel, ctx);
break;
@@ -3,6 +3,8 @@
import type {AttachmentID, ChannelID, EmojiID, GuildID, MessageID, UserID} from '@app/api/BrandedTypes';
import {IChannelRepository} from '@app/api/channel/IChannelRepository';
import {ChannelRepository as NewChannelRepository} from '@app/api/channel/repositories/ChannelRepository';
import type {GuildChannelListMode} from '@app/api/channel/repositories/IChannelDataRepository';
import type {UpsertMessageOptions} from '@app/api/channel/repositories/IMessageRepository';
import type {ChannelRow} from '@app/api/database/types/ChannelTypes';
import type {MessageRow} from '@app/api/database/types/MessageTypes';
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
@@ -30,6 +32,10 @@ export class ChannelRepository extends IChannelRepository {
return this.repository.messageInteractions;
}
get threads() {
return this.repository.threads;
}
get crossposts() {
return this.repository.crossposts;
}
@@ -42,12 +48,12 @@ export class ChannelRepository extends IChannelRepository {
return this.repository.channelData.upsert(data);
}
async updateLastMessageId(channelId: ChannelID, messageId: MessageID): Promise<void> {
return this.repository.channelData.updateLastMessageId(channelId, messageId);
async updateLastMessageId(channelId: ChannelID, messageId: MessageID, opts?: {isInsert?: boolean}): Promise<void> {
return this.repository.channelData.updateLastMessageId(channelId, messageId, opts);
}
async delete(channelId: ChannelID, guildId?: GuildID): Promise<void> {
return this.repository.channelData.delete(channelId, guildId);
async delete(channelId: ChannelID, guildId?: GuildID, type?: number): Promise<void> {
return this.repository.channelData.delete(channelId, guildId, type);
}
async listMessages(
@@ -63,8 +69,8 @@ export class ChannelRepository extends IChannelRepository {
return this.repository.messages.getMessage(channelId, messageId);
}
async upsertMessage(data: MessageRow, oldData?: MessageRow | null): Promise<Message> {
return this.repository.messages.upsertMessage(data, oldData);
async upsertMessage(data: MessageRow, oldData?: MessageRow | null, opts?: UpsertMessageOptions): Promise<Message> {
return this.repository.messages.upsertMessage(data, oldData, opts);
}
async deleteMessage(
@@ -176,8 +182,8 @@ export class ChannelRepository extends IChannelRepository {
);
}
async listGuildChannels(guildId: GuildID): Promise<Array<Channel>> {
return this.repository.channelData.listGuildChannels(guildId);
async listGuildChannels(guildId: GuildID, mode: GuildChannelListMode): Promise<Array<Channel>> {
return this.repository.channelData.listGuildChannels(guildId, mode);
}
async listChannels(channelIds: Array<ChannelID>): Promise<Array<Channel>> {
@@ -1,7 +1,9 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {AttachmentID, ChannelID, EmojiID, GuildID, MessageID, UserID} from '@app/api/BrandedTypes';
import type {GuildChannelListMode} from '@app/api/channel/repositories/IChannelDataRepository';
import {IChannelRepositoryAggregate} from '@app/api/channel/repositories/IChannelRepositoryAggregate';
import type {UpsertMessageOptions} from '@app/api/channel/repositories/IMessageRepository';
import type {ChannelRow} from '@app/api/database/types/ChannelTypes';
import type {MessageRow} from '@app/api/database/types/MessageTypes';
import type {Channel} from '@app/api/models/Channel';
@@ -13,11 +15,11 @@ export abstract class IChannelRepository extends IChannelRepositoryAggregate {
abstract upsert(data: ChannelRow): Promise<Channel>;
abstract updateLastMessageId(channelId: ChannelID, messageId: MessageID): Promise<void>;
abstract updateLastMessageId(channelId: ChannelID, messageId: MessageID, opts?: {isInsert?: boolean}): Promise<void>;
abstract delete(channelId: ChannelID, guildId?: GuildID): Promise<void>;
abstract delete(channelId: ChannelID, guildId?: GuildID, type?: number): Promise<void>;
abstract listGuildChannels(guildId: GuildID): Promise<Array<Channel>>;
abstract listGuildChannels(guildId: GuildID, mode: GuildChannelListMode): Promise<Array<Channel>>;
abstract listChannels(channelIds: Array<ChannelID>): Promise<Array<Channel>>;
@@ -32,7 +34,7 @@ export abstract class IChannelRepository extends IChannelRepositoryAggregate {
abstract getMessage(channelId: ChannelID, messageId: MessageID): Promise<Message | null>;
abstract upsertMessage(data: MessageRow, oldData?: MessageRow | null): Promise<Message>;
abstract upsertMessage(data: MessageRow, oldData?: MessageRow | null, opts?: UpsertMessageOptions): Promise<Message>;
abstract deleteMessage(
channelId: ChannelID,
@@ -2,6 +2,8 @@
import {requireSudoMode} from '@app/api/auth/services/SudoVerificationService';
import {createChannelID, createUserID} from '@app/api/BrandedTypes';
import {GatedJsonValidator} from '@app/api/channel/threads/GatedJsonValidator';
import {viewerActive, viewerFromCtx} from '@app/api/experiment/ChannelThreadsGate';
import {DefaultUserOnly, LoginRequired} from '@app/api/middleware/AuthMiddleware';
import {GroupDmRecipientAddProtectionMiddleware} from '@app/api/middleware/GroupDmProtectionMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
@@ -9,13 +11,14 @@ import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
import {SudoModeMiddleware} from '@app/api/middleware/SudoModeMiddleware';
import {RateLimitConfigs} from '@app/api/RateLimitConfig';
import type {HonoApp, HonoEnv} from '@app/api/types/HonoEnv';
import {CLIENT_FEATURES_HEADER, parseClientFeaturesHeader} from '@app/api/utils/featureUtils';
import {Validator} from '@app/api/Validator';
import {ANNOUNCEMENT_CONVERTIBLE_CHANNEL_TYPES} from '@fluxer/constants/src/ChannelConstants';
import {TEXT_THREAD_PARENT_CHANNEL_TYPES, THREAD_ONLY_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {ChannelTypeConversionNotSupportedError} from '@fluxer/errors/src/domains/channel/ChannelTypeConversionNotSupportedError';
import {UnknownChannelError} from '@fluxer/errors/src/domains/channel/UnknownChannelError';
import {SudoVerificationSchema} from '@fluxer/schema/src/domains/auth/AuthSchemas';
import {
ChannelUpdateGatedRequest,
ChannelUpdateRequest,
ChannelUpdateRequestBody,
DeleteChannelQuery,
@@ -33,6 +36,8 @@ import {
} from '@fluxer/schema/src/domains/common/CommonParamSchemas';
import type {Context} from 'hono';
const THREAD_PARENT_DEFAULT_KEYS = ['default_auto_archive_duration', 'default_thread_rate_limit_per_user'];
function isPlainObject(value: unknown): value is Record<string, unknown> {
return typeof value === 'object' && value !== null && !Array.isArray(value);
}
@@ -59,11 +64,7 @@ export function ChannelController(app: HonoApp) {
const requestCache = ctx.get('requestCache');
const channelRequestService = ctx.get('channelRequestService');
return ctx.json(
await channelRequestService.getChannelResponse({
userId,
channelId,
requestCache,
}),
await channelRequestService.getChannelResponse({viewer: viewerFromCtx(ctx), userId, channelId, requestCache}),
);
},
);
@@ -86,7 +87,7 @@ export function ChannelController(app: HonoApp) {
const user = ctx.get('user');
const channelId = createChannelID(ctx.req.valid('param').channel_id);
const channelRequestService = ctx.get('channelRequestService');
return ctx.json(await channelRequestService.getSlowmodeState({user, channelId}));
return ctx.json(await channelRequestService.getSlowmodeState({viewer: viewerFromCtx(ctx), user, channelId}));
},
);
app.get(
@@ -109,7 +110,7 @@ export function ChannelController(app: HonoApp) {
const userId = ctx.get('user').id;
const channelId = createChannelID(ctx.req.valid('param').channel_id);
const channelRequestService = ctx.get('channelRequestService');
return ctx.json(await channelRequestService.listRtcRegions({userId, channelId}));
return ctx.json(await channelRequestService.listRtcRegions({viewer: viewerFromCtx(ctx), userId, channelId}));
},
);
app.patch(
@@ -123,15 +124,17 @@ export function ChannelController(app: HonoApp) {
}
const channelId = createChannelID(result.data.channel_id);
const existing = await ctx.get('channelService').channelData.operations.getChannel({
viewer: viewerFromCtx(ctx),
userId: ctx.get('user').id,
channelId,
skipNsfwValidation: true,
});
ctx.set('channelUpdateType', existing.type);
ctx.set('channelUpdateGuildId', existing.guildId?.toString());
return undefined;
},
}),
Validator('json', ChannelUpdateRequest, {
GatedJsonValidator(ChannelUpdateRequest, ChannelUpdateGatedRequest, {
pre: async (raw: unknown, ctx: Context<HonoEnv>) => {
const channelType = ctx.get('channelUpdateType');
if (channelType === undefined) {
@@ -152,6 +155,20 @@ export function ChannelController(app: HonoApp) {
}
return {...body, type: requestedType};
},
touchesGate: (body, ctx) => {
const channelType = ctx.get('channelUpdateType');
if (channelType === undefined) return false;
if (THREAD_ONLY_CHANNEL_TYPES.has(channelType)) return true;
return (
TEXT_THREAD_PARENT_CHANNEL_TYPES.has(channelType) &&
isPlainObject(body) &&
THREAD_PARENT_DEFAULT_KEYS.some((key) => body[key] !== undefined)
);
},
active: (ctx) => {
const guildId = ctx.get('channelUpdateGuildId');
return guildId !== undefined && viewerActive(viewerFromCtx(ctx), guildId);
},
}),
OpenAPI({
operationId: 'update_channel',
@@ -171,12 +188,13 @@ export function ChannelController(app: HonoApp) {
const existingType = ctx.get('channelUpdateType');
const typeConversion =
existingType !== undefined && data.type !== existingType ? {from: existingType, to: data.type} : null;
const clientFeatures = parseClientFeaturesHeader(ctx.req.header(CLIENT_FEATURES_HEADER));
const clientFeatures = ctx.get('clientFeatures');
const requestCache = ctx.get('requestCache');
const auditLogReason = ctx.get('auditLogReason') ?? null;
const channelRequestService = ctx.get('channelRequestService');
return ctx.json(
await channelRequestService.updateChannel({
viewer: viewerFromCtx(ctx),
userId,
channelId,
data,
@@ -217,14 +235,23 @@ export function ChannelController(app: HonoApp) {
const requestCache = ctx.get('requestCache');
const auditLogReason = ctx.get('auditLogReason') ?? null;
const channelRequestService = ctx.get('channelRequestService');
await ctx.get('channelService').channelData.operations.getChannel({userId, channelId});
await ctx
.get('channelService')
.channelData.operations.getChannel({viewer: viewerFromCtx(ctx), userId, channelId});
if (delete_messages) {
await requireSudoMode(ctx, user, body);
await ctx.get('channelService').userMessageDeletion.deleteUserMessagesInScope(userId, {
channelIds: [channelId],
});
}
await channelRequestService.deleteChannel({userId, channelId, requestCache, silent, auditLogReason});
await channelRequestService.deleteChannel({
viewer: viewerFromCtx(ctx),
userId,
channelId,
requestCache,
silent,
auditLogReason,
});
return ctx.body(null, 204);
},
);
@@ -286,7 +313,9 @@ export function ChannelController(app: HonoApp) {
const body = ctx.req.valid('json');
const requestCache = ctx.get('requestCache');
if (delete_messages && recipientId === userId) {
await ctx.get('channelService').channelData.operations.getChannel({userId, channelId});
await ctx
.get('channelService')
.channelData.operations.getChannel({viewer: viewerFromCtx(ctx), userId, channelId});
await requireSudoMode(ctx, ctx.get('user'), body);
await ctx.get('channelService').userMessageDeletion.deleteUserMessagesInScope(userId, {
channelIds: [channelId],
@@ -319,7 +348,7 @@ export function ChannelController(app: HonoApp) {
const channelId = createChannelID(ctx.req.valid('param').channel_id);
const overwriteId = ctx.req.valid('param').overwrite_id;
const data = ctx.req.valid('json');
const clientFeatures = parseClientFeaturesHeader(ctx.req.header(CLIENT_FEATURES_HEADER));
const clientFeatures = ctx.get('clientFeatures');
const requestCache = ctx.get('requestCache');
const auditLogReason = ctx.get('auditLogReason') ?? null;
await ctx.get('channelService').channelData.operations.setChannelPermissionOverwrite({
@@ -332,6 +361,7 @@ export function ChannelController(app: HonoApp) {
deny_: data.deny ? data.deny : 0n,
},
clientFeatures,
viewer: viewerFromCtx(ctx),
requestCache,
auditLogReason,
});
@@ -363,6 +393,8 @@ export function ChannelController(app: HonoApp) {
userId,
channelId,
overwriteId,
clientFeatures: ctx.get('clientFeatures'),
viewer: viewerFromCtx(ctx),
requestCache,
auditLogReason,
});
@@ -1,6 +1,7 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {createChannelID} from '@app/api/BrandedTypes';
import {viewerFromCtx} from '@app/api/experiment/ChannelThreadsGate';
import {LoginRequired} from '@app/api/middleware/AuthMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
@@ -37,6 +38,7 @@ export function ChannelFollowController(app: HonoApp) {
assertAccountNotLimited(ctx.get('user'));
const followed = await ctx.get('channelFollowService').followChannel({
userId: ctx.get('user').id,
viewer: viewerFromCtx(ctx),
channelId: createChannelID(ctx.req.valid('param').channel_id),
webhookChannelId: createChannelID(ctx.req.valid('json').webhook_channel_id),
requestCache: ctx.get('requestCache'),
@@ -66,6 +68,7 @@ export function ChannelFollowController(app: HonoApp) {
async (ctx) => {
const stats = await ctx.get('channelFollowService').getFollowerStats({
userId: ctx.get('user').id,
viewer: viewerFromCtx(ctx),
channelId: createChannelID(ctx.req.valid('param').channel_id),
});
return ctx.json(stats);
@@ -0,0 +1,184 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {createChannelID} from '@app/api/BrandedTypes';
import {ChannelThreadsRouteGuard} from '@app/api/channel/threads/ChannelThreadsRouteGuard';
import {viewerFromCtx} from '@app/api/experiment/ChannelThreadsGate';
import {LoginRequired} from '@app/api/middleware/AuthMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
import {RateLimitConfigs} from '@app/api/RateLimitConfig';
import type {HonoApp, HonoEnv} from '@app/api/types/HonoEnv';
import {Validator} from '@app/api/Validator';
import {ChannelResponse} from '@fluxer/schema/src/domains/channel/ChannelSchemas';
import {
ChannelIdTagIdParam,
ForumTagRequest,
SearchIndexNotReadyResponse,
ThreadPostDataRequest,
ThreadPostDataResponse,
ThreadSearchQuery,
ThreadSearchResult,
} from '@fluxer/schema/src/domains/channel/ForumRequestSchemas';
import {ChannelIdParam} from '@fluxer/schema/src/domains/common/CommonParamSchemas';
import type {Context} from 'hono';
const EXPERIMENT = 'channel_threads';
const TAGS = 'Channels';
function tagEditParams(ctx: Context<HonoEnv>) {
return {
userId: ctx.get('user').id,
viewer: viewerFromCtx(ctx),
clientFeatures: ctx.get('clientFeatures'),
requestCache: ctx.get('requestCache'),
auditLogReason: ctx.get('auditLogReason') ?? null,
};
}
export function ForumController(app: HonoApp) {
app.get(
'/channels/:channel_id/threads/search',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREADS_SEARCH),
LoginRequired,
Validator('param', ChannelIdParam),
Validator('query', ThreadSearchQuery),
OpenAPI({
operationId: 'search_threads',
summary: 'Search threads',
description:
'Returns threads of the channel that match the search. Requires the read message history permission. While the search index of the guild is being built, responds with 202 and a body that says when to retry.',
responseSchema: ThreadSearchResult,
acceptedResponseSchema: SearchIndexNotReadyResponse,
statusCode: [200, 202],
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
return ctx.json(
await ctx.get('threadService').forum.search({
viewer: viewerFromCtx(ctx),
userId: ctx.get('user').id,
channelId: createChannelID(ctx.req.valid('param').channel_id),
query: ctx.req.valid('query'),
}),
);
},
);
app.post(
'/channels/:channel_id/post-data',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_POST_DATA),
LoginRequired,
Validator('param', ChannelIdParam),
Validator('json', ThreadPostDataRequest),
OpenAPI({
operationId: 'get_channel_post_data',
summary: 'Get forum post data',
description:
'Returns the owner and first message of each requested post in a forum or media channel. Requires the read message history permission.',
responseSchema: ThreadPostDataResponse,
statusCode: 200,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
return ctx.json(
await ctx.get('threadService').forum.postData({
viewer: viewerFromCtx(ctx),
userId: ctx.get('user').id,
channelId: createChannelID(ctx.req.valid('param').channel_id),
threadIds: ctx.req.valid('json').thread_ids.map((id) => createChannelID(id)),
requestCache: ctx.get('requestCache'),
}),
);
},
);
app.post(
'/channels/:channel_id/tags',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_FORUM_TAGS),
LoginRequired,
Validator('param', ChannelIdParam),
Validator('json', ForumTagRequest),
OpenAPI({
operationId: 'create_forum_tag',
summary: 'Create a forum tag',
description:
'Adds a tag to a forum or media channel. Requires the manage channels permission. Returns the updated channel.',
responseSchema: ChannelResponse,
statusCode: 200,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
return ctx.json(
await ctx.get('channelRequestService').editForumTags({
...tagEditParams(ctx),
channelId: createChannelID(ctx.req.valid('param').channel_id),
edit: {kind: 'create', tag: ctx.req.valid('json')},
}),
);
},
);
app.put(
'/channels/:channel_id/tags/:tag_id',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_FORUM_TAGS),
LoginRequired,
Validator('param', ChannelIdTagIdParam),
Validator('json', ForumTagRequest),
OpenAPI({
operationId: 'update_forum_tag',
summary: 'Update a forum tag',
description:
'Replaces a tag of a forum or media channel. Requires the manage channels permission. Returns the updated channel.',
responseSchema: ChannelResponse,
statusCode: 200,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {channel_id, tag_id} = ctx.req.valid('param');
return ctx.json(
await ctx.get('channelRequestService').editForumTags({
...tagEditParams(ctx),
channelId: createChannelID(channel_id),
edit: {kind: 'update', tagId: tag_id, tag: ctx.req.valid('json')},
}),
);
},
);
app.delete(
'/channels/:channel_id/tags/:tag_id',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_FORUM_TAGS),
LoginRequired,
Validator('param', ChannelIdTagIdParam),
OpenAPI({
operationId: 'delete_forum_tag',
summary: 'Delete a forum tag',
description:
'Removes a tag from a forum or media channel. Requires the manage channels permission. Returns the updated channel.',
responseSchema: ChannelResponse,
statusCode: 200,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {channel_id, tag_id} = ctx.req.valid('param');
return ctx.json(
await ctx.get('channelRequestService').editForumTags({
...tagEditParams(ctx),
channelId: createChannelID(channel_id),
edit: {kind: 'delete', tagId: tag_id},
}),
);
},
);
}
@@ -6,11 +6,13 @@ import {Config} from '@app/api/Config';
import type {MessageRequest, MessageUpdateRequest} from '@app/api/channel/MessageTypes';
import {normalizeMessageRequestPayload} from '@app/api/channel/services/message/MessageRequestCompatibility';
import {parseMultipartMessageData} from '@app/api/channel/services/message/MessageRequestParser';
import {viewerFromCtx} from '@app/api/experiment/ChannelThreadsGate';
import {DefaultUserOnly, LoginRequired} from '@app/api/middleware/AuthMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
import {SudoModeMiddleware} from '@app/api/middleware/SudoModeMiddleware';
import {RateLimitConfigs} from '@app/api/RateLimitConfig';
import {readStateCapable} from '@app/api/read_state/ReadStateChannelMeta';
import type {HonoApp} from '@app/api/types/HonoEnv';
import {assertAccountNotLimited} from '@app/api/user/AccountLimit';
import {parseJsonPreservingLargeIntegers} from '@app/api/utils/LosslessJsonParser';
@@ -71,6 +73,7 @@ export function MessageController(app: HonoApp) {
const messageRequestService = ctx.get('messageRequestService');
return ctx.json(
await messageRequestService.listMessages({
viewer: viewerFromCtx(ctx),
userId,
channelId,
query: {
@@ -108,6 +111,7 @@ export function MessageController(app: HonoApp) {
const messageRequestService = ctx.get('messageRequestService');
return ctx.json(
await messageRequestService.listMessagesBulk({
viewer: viewerFromCtx(ctx),
userId,
requests: requests.map((request) => ({
channelId: createChannelID(request.channel_id),
@@ -148,6 +152,7 @@ export function MessageController(app: HonoApp) {
const messageRequestService = ctx.get('messageRequestService');
return ctx.json(
await messageRequestService.getMessage({
viewer: viewerFromCtx(ctx),
userId,
channelId,
messageId,
@@ -196,6 +201,7 @@ export function MessageController(app: HonoApp) {
return validationResult.data;
})();
const response = await messageRequestService.sendMessage({
viewer: viewerFromCtx(ctx),
user,
channelId,
data: validatedData as MessageRequest,
@@ -229,6 +235,7 @@ export function MessageController(app: HonoApp) {
const {attachments} = ctx.req.valid('json');
return ctx.json({
attachments: await ctx.get('channelService').attachments.requestPresignedAttachmentUploadUrls({
viewer: viewerFromCtx(ctx),
userId: ctx.get('user').id,
channelId,
clientIp,
@@ -262,6 +269,7 @@ export function MessageController(app: HonoApp) {
const {uploads} = ctx.req.valid('json');
return ctx.json({
uploads: await ctx.get('channelService').attachments.completeMultipartAttachmentUploads({
viewer: viewerFromCtx(ctx),
userId: ctx.get('user').id,
channelId,
clientIp,
@@ -314,6 +322,7 @@ export function MessageController(app: HonoApp) {
})();
return ctx.json(
await messageRequestService.editMessage({
viewer: viewerFromCtx(ctx),
userId,
channelId,
messageId,
@@ -364,9 +373,14 @@ export function MessageController(app: HonoApp) {
const messageId = createMessageID(message_id);
const requestCache = ctx.get('requestCache');
const auditLogReason = ctx.get('auditLogReason') ?? null;
await ctx
.get('channelService')
.messages.deletion.deleteMessage({userId, channelId, messageId, requestCache, auditLogReason});
await ctx.get('channelService').messages.deletion.deleteMessage({
viewer: viewerFromCtx(ctx),
userId,
channelId,
messageId,
requestCache,
auditLogReason,
});
return ctx.body(null, 204);
},
);
@@ -393,6 +407,7 @@ export function MessageController(app: HonoApp) {
const attachmentId = createAttachmentID(attachment_id);
const requestCache = ctx.get('requestCache');
await ctx.get('channelService').attachments.deleteAttachment({
viewer: viewerFromCtx(ctx),
userId,
channelId,
messageId: messageId,
@@ -423,9 +438,13 @@ export function MessageController(app: HonoApp) {
const channelId = createChannelID(ctx.req.valid('param').channel_id);
const messageIds = ctx.req.valid('json').message_ids.map(createMessageID);
const auditLogReason = ctx.get('auditLogReason') ?? null;
await ctx
.get('channelService')
.messages.deletion.bulkDeleteMessages({userId, channelId, messageIds, auditLogReason});
await ctx.get('channelService').messages.deletion.bulkDeleteMessages({
viewer: viewerFromCtx(ctx),
userId,
channelId,
messageIds,
auditLogReason,
});
return ctx.body(null, 204);
},
);
@@ -449,7 +468,7 @@ export function MessageController(app: HonoApp) {
const channelId = createChannelID(ctx.req.valid('param').channel_id);
const {deletedCount} = await ctx
.get('channelService')
.messages.deletion.purgePersonalNotesMessages({userId, channelId});
.messages.deletion.purgePersonalNotesMessages({viewer: viewerFromCtx(ctx), userId, channelId});
return ctx.json({deleted_count: deletedCount});
},
);
@@ -475,7 +494,9 @@ export function MessageController(app: HonoApp) {
const userId = user.id;
const channelId = createChannelID(ctx.req.valid('param').channel_id);
const body = ctx.req.valid('json');
await ctx.get('channelService').channelData.operations.getChannel({userId, channelId});
await ctx
.get('channelService')
.channelData.operations.getChannel({viewer: viewerFromCtx(ctx), userId, channelId});
await requireSudoMode(ctx, user, body);
await ctx.get('channelService').userMessageDeletion.deleteUserMessagesInScope(userId, {
channelIds: [channelId],
@@ -492,17 +513,20 @@ export function MessageController(app: HonoApp) {
operationId: 'indicate_typing',
summary: 'Indicate typing activity',
responseSchema: null,
statusCode: 204,
statusCode: [200, 204],
security: ['botToken', 'bearerToken', 'sessionToken'],
tags: ['Channels', 'Messages'],
description:
'Notifies other users in the channel that you are actively typing. Typing indicators typically expire after a short period (usually 10 seconds). Returns 204 No Content. Commonly called repeatedly while the user is composing a message.',
'Notifies other users in the channel that you are actively typing. Typing indicators typically expire after a short period (usually 10 seconds). Returns 204 No Content, or 200 with a JSON body holding the remaining slowmode cooldowns in message_send_cooldown_ms and thread_create_cooldown_ms when the user is rate limited. Commonly called repeatedly while the user is composing a message.',
}),
async (ctx) => {
const userId = ctx.get('user').id;
const user = ctx.get('user');
const viewer = viewerFromCtx(ctx);
const channelId = createChannelID(ctx.req.valid('param').channel_id);
await ctx.get('channelService').interactions.startTyping({userId, channelId});
return ctx.body(null, 204);
const channelService = ctx.get('channelService');
const auth = await channelService.interactions.startTyping({viewer, userId: user.id, channelId});
const cooldown = await channelService.getTypingCooldown({user, viewer, auth});
return cooldown ? ctx.json(cooldown, 200) : ctx.body(null, 204);
},
);
app.post(
@@ -526,6 +550,7 @@ export function MessageController(app: HonoApp) {
return ctx.json(
await ctx.get('messageRequestService').crosspostMessage({
userId: ctx.get('user').id,
viewer: viewerFromCtx(ctx),
channelId: createChannelID(channel_id),
messageId: createMessageID(message_id),
requestCache: ctx.get('requestCache'),
@@ -553,6 +578,7 @@ export function MessageController(app: HonoApp) {
return ctx.json(
await ctx.get('messageRequestService').getCrosspostSource({
userId: ctx.get('user').id,
viewer: viewerFromCtx(ctx),
channelId: createChannelID(channel_id),
messageId: createMessageID(message_id),
requestCache: ctx.get('requestCache'),
@@ -589,6 +615,7 @@ export function MessageController(app: HonoApp) {
messageId,
mentionCount: mentionCount ?? 0,
manual,
capable: readStateCapable(ctx),
});
return ctx.body(null, 204);
},
@@ -3,14 +3,17 @@
import {createChannelID, createMessageID, createUserID} from '@app/api/BrandedTypes';
import {isPersonalNotesChannel} from '@app/api/channel/services/message/MessageHelpers';
import {SYSTEM_USER_ID} from '@app/api/constants/Core';
import {THREAD_FEATURE_CHANNEL_TYPES, viewerActive, viewerFromCtx} from '@app/api/experiment/ChannelThreadsGate';
import {LoginRequired} from '@app/api/middleware/AuthMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
import {RateLimitConfigs} from '@app/api/RateLimitConfig';
import {readStateCapable} from '@app/api/read_state/ReadStateChannelMeta';
import type {HonoApp} from '@app/api/types/HonoEnv';
import {Validator} from '@app/api/Validator';
import {UserFlags} from '@fluxer/constants/src/UserConstants';
import {ValidationErrorCodes} from '@fluxer/constants/src/ValidationErrorCodes';
import {InvalidChannelTypeError} from '@fluxer/errors/src/domains/channel/InvalidChannelTypeError';
import {UnclaimedAccountCannotAddReactionsError} from '@fluxer/errors/src/domains/channel/UnclaimedAccountCannotAddReactionsError';
import {InputValidationError} from '@fluxer/errors/src/domains/core/InputValidationError';
import {
@@ -53,9 +56,14 @@ export function MessageInteractionController(app: HonoApp) {
const requestCache = ctx.get('requestCache');
const {limit, before} = ctx.req.valid('query');
return ctx.json(
await ctx
.get('channelService')
.interactions.getChannelPins({userId, channelId, requestCache, limit, beforeTimestamp: before}),
await ctx.get('channelService').interactions.getChannelPins({
viewer: viewerFromCtx(ctx),
userId,
channelId,
requestCache,
limit,
beforeTimestamp: before,
}),
);
},
);
@@ -77,10 +85,32 @@ export function MessageInteractionController(app: HonoApp) {
async (ctx) => {
const userId = ctx.get('user').id;
const channelId = createChannelID(ctx.req.valid('param').channel_id);
const channel = await ctx.get('channelService').channelData.operations.getChannelSystem(channelId);
const channelService = ctx.get('channelService');
const channel = await channelService.channelData.operations.getChannelSystem(channelId);
if (channel && THREAD_FEATURE_CHANNEL_TYPES.has(channel.type)) {
const viewer = viewerFromCtx(ctx);
if (channel.guildId === null || !viewerActive(viewer, channel.guildId)) {
return ctx.body(null, 204);
}
if (channel.isThreadOnly()) {
throw new InvalidChannelTypeError();
}
await channelService.channelData.auth.getChannelAuthenticated({
userId,
channelId,
viewer,
skipNsfwValidation: true,
});
}
const timestamp = channel?.lastPinTimestamp;
if (timestamp != null) {
await ctx.get('readStateService').ackPins({userId, channelId, timestamp});
await ctx.get('readStateService').ackPins({
userId,
channelId,
timestamp,
capable: readStateCapable(ctx),
channel,
});
}
return ctx.body(null, 204);
},
@@ -108,6 +138,7 @@ export function MessageInteractionController(app: HonoApp) {
const requestCache = ctx.get('requestCache');
const auditLogReason = ctx.get('auditLogReason') ?? null;
await ctx.get('channelService').interactions.pinMessage({
viewer: viewerFromCtx(ctx),
userId,
channelId,
messageId,
@@ -140,6 +171,7 @@ export function MessageInteractionController(app: HonoApp) {
const requestCache = ctx.get('requestCache');
const auditLogReason = ctx.get('auditLogReason') ?? null;
await ctx.get('channelService').interactions.unpinMessage({
viewer: viewerFromCtx(ctx),
userId,
channelId,
messageId,
@@ -172,9 +204,15 @@ export function MessageInteractionController(app: HonoApp) {
const channelId = createChannelID(channel_id);
const messageId = createMessageID(message_id);
const afterUserId = after ? createUserID(after) : undefined;
const result = await ctx
.get('channelService')
.interactions.getUsersForReaction({userId, channelId, messageId, emoji, limit, after: afterUserId});
const result = await ctx.get('channelService').interactions.getUsersForReaction({
viewer: viewerFromCtx(ctx),
userId,
channelId,
messageId,
emoji,
limit,
after: afterUserId,
});
ctx.header('X-Has-More', result.has_more ? 'true' : 'false');
if (result.next_after !== null) {
ctx.header('X-Next-After', result.next_after);
@@ -205,9 +243,15 @@ export function MessageInteractionController(app: HonoApp) {
const channelId = createChannelID(channel_id);
const messageId = createMessageID(message_id);
const afterUserId = after ? createUserID(after) : undefined;
const result = await ctx
.get('channelService')
.interactions.getUsersForReaction({userId, channelId, messageId, emoji, limit, after: afterUserId});
const result = await ctx.get('channelService').interactions.getUsersForReaction({
viewer: viewerFromCtx(ctx),
userId,
channelId,
messageId,
emoji,
limit,
after: afterUserId,
});
return ctx.json(
{
items: result.users,
@@ -248,6 +292,7 @@ export function MessageInteractionController(app: HonoApp) {
throw InputValidationError.fromCode('emoji', ValidationErrorCodes.MUST_START_SESSION_BEFORE_SENDING);
}
await ctx.get('channelService').interactions.addReaction({
viewer: viewerFromCtx(ctx),
userId: user.id,
sessionId,
channelId,
@@ -282,6 +327,7 @@ export function MessageInteractionController(app: HonoApp) {
const sessionId = ctx.req.valid('query').session_id;
const requestCache = ctx.get('requestCache');
await ctx.get('channelService').interactions.removeOwnReaction({
viewer: viewerFromCtx(ctx),
userId,
sessionId,
channelId,
@@ -317,6 +363,7 @@ export function MessageInteractionController(app: HonoApp) {
const sessionId = ctx.req.valid('query').session_id;
const requestCache = ctx.get('requestCache');
await ctx.get('channelService').interactions.removeReaction({
viewer: viewerFromCtx(ctx),
userId,
sessionId,
channelId,
@@ -348,7 +395,9 @@ export function MessageInteractionController(app: HonoApp) {
const userId = ctx.get('user').id;
const channelId = createChannelID(channel_id);
const messageId = createMessageID(message_id);
await ctx.get('channelService').interactions.removeAllReactionsForEmoji({userId, channelId, messageId, emoji});
await ctx
.get('channelService')
.interactions.removeAllReactionsForEmoji({viewer: viewerFromCtx(ctx), userId, channelId, messageId, emoji});
return ctx.body(null, 204);
},
);
@@ -372,7 +421,9 @@ export function MessageInteractionController(app: HonoApp) {
const userId = ctx.get('user').id;
const channelId = createChannelID(channel_id);
const messageId = createMessageID(message_id);
await ctx.get('channelService').interactions.removeAllReactions({userId, channelId, messageId});
await ctx
.get('channelService')
.interactions.removeAllReactions({viewer: viewerFromCtx(ctx), userId, channelId, messageId});
return ctx.body(null, 204);
},
);
@@ -2,6 +2,7 @@
import {createChannelID} from '@app/api/BrandedTypes';
import {Config} from '@app/api/Config';
import {viewerFromCtx} from '@app/api/experiment/ChannelThreadsGate';
import {DefaultUserOnly, LoginRequired} from '@app/api/middleware/AuthMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
@@ -40,7 +41,9 @@ export function StreamController(app: HonoApp) {
const user = ctx.get('user');
const {region} = ctx.req.valid('json');
const streamKey = ctx.req.valid('param').stream_key;
await ctx.get('streamService').updateStreamRegion({streamKey, region, userId: user.id});
await ctx
.get('streamService')
.updateStreamRegion({viewer: viewerFromCtx(ctx), streamKey, region, userId: user.id});
return ctx.body(null, 204);
},
);
@@ -64,7 +67,9 @@ export function StreamController(app: HonoApp) {
async (ctx) => {
const user = ctx.get('user');
const streamKey = ctx.req.valid('param').stream_key;
const preview = await ctx.get('streamService').getPreview({streamKey, userId: user.id});
const preview = await ctx
.get('streamService')
.getPreview({viewer: viewerFromCtx(ctx), streamKey, userId: user.id});
if (!preview) {
return ctx.body(null, 404);
}
@@ -102,6 +107,7 @@ export function StreamController(app: HonoApp) {
clientIpHeaderName: Config.proxy.client_ip_header,
});
const response = await ctx.get('streamService').createPreviewUploadUrl({
viewer: viewerFromCtx(ctx),
streamKey,
channelId: createChannelID(channel_id),
userId: user.id,
@@ -133,6 +139,7 @@ export function StreamController(app: HonoApp) {
const {thumbnail, channel_id, content_type} = ctx.req.valid('json');
const streamKey = ctx.req.valid('param').stream_key;
await ctx.get('streamService').uploadPreview({
viewer: viewerFromCtx(ctx),
streamKey,
channelId: createChannelID(channel_id),
userId: user.id,
@@ -161,7 +168,7 @@ export function StreamController(app: HonoApp) {
async (ctx) => {
const user = ctx.get('user');
const streamKey = ctx.req.valid('param').stream_key;
await ctx.get('streamService').deletePreview({streamKey, userId: user.id});
await ctx.get('streamService').deletePreview({viewer: viewerFromCtx(ctx), streamKey, userId: user.id});
return ctx.body(null, 204);
},
);
@@ -0,0 +1,530 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {type ChannelID, createChannelID, createGuildID, createMessageID, createUserID} from '@app/api/BrandedTypes';
import type {MessageRequest} from '@app/api/channel/MessageTypes';
import {normalizeMessageRequestPayload} from '@app/api/channel/services/message/MessageRequestCompatibility';
import {parseMultipartMessageData} from '@app/api/channel/services/message/MessageRequestParser';
import type {ForumPostInput} from '@app/api/channel/services/thread/ThreadCreationService';
import {ChannelThreadsRouteGuard} from '@app/api/channel/threads/ChannelThreadsRouteGuard';
import {viewerFromCtx} from '@app/api/experiment/ChannelThreadsGate';
import {BotOnly, LoginRequired} from '@app/api/middleware/AuthMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
import {RateLimitConfigs} from '@app/api/RateLimitConfig';
import type {HonoApp, HonoEnv} from '@app/api/types/HonoEnv';
import {parseJsonPreservingLargeIntegers} from '@app/api/utils/LosslessJsonParser';
import {inputValidationErrorFromZodIssues, Validator} from '@app/api/Validator';
import {TEXT_THREAD_PARENT_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {ValidationErrorCodes} from '@fluxer/constants/src/ValidationErrorCodes';
import {InvalidChannelTypeError} from '@fluxer/errors/src/domains/channel/InvalidChannelTypeError';
import {InputValidationError} from '@fluxer/errors/src/domains/core/InputValidationError';
import {
ForumThreadMessageRequest,
StartForumThreadRequest,
StartForumThreadResponse,
StartThreadMultipartRequest,
StartThreadRequestBody,
} from '@fluxer/schema/src/domains/channel/ForumRequestSchemas';
import {
ActiveThreadsResponse,
ArchivedThreadsQuery,
ArchivedThreadsResponse,
JoinedArchivedThreadsQuery,
StartThreadFromMessageRequest,
StartThreadRequest,
ThreadChannelResponse,
ThreadLocationQuery,
ThreadMemberGetQuery,
ThreadMemberListResponse,
ThreadMembersListQuery,
} from '@fluxer/schema/src/domains/channel/ThreadRequestSchemas';
import {ThreadMemberResponse} from '@fluxer/schema/src/domains/channel/ThreadSchemas';
import {
ChannelIdMessageIdParam,
ChannelIdParam,
ChannelIdUserIdParam,
GuildIdParam,
} from '@fluxer/schema/src/domains/common/CommonParamSchemas';
import type {Context, MiddlewareHandler} from 'hono';
import {z} from 'zod';
const EXPERIMENT = 'channel_threads';
const TAGS = 'Channels';
function isMultipart(ctx: Context<HonoEnv>): boolean {
return (ctx.req.header('content-type') ?? '').includes('multipart/form-data');
}
async function readPayloadJson(ctx: Context<HonoEnv>): Promise<unknown> {
let payloadJson: unknown;
try {
payloadJson = (await ctx.req.parseBody())['payload_json'];
} catch {
throw InputValidationError.fromCode('multipart_form', ValidationErrorCodes.FAILED_TO_PARSE_MULTIPART_FORM_DATA);
}
if (payloadJson === undefined) return {};
if (typeof payloadJson === 'string') {
try {
return parseJsonPreservingLargeIntegers(payloadJson);
} catch {}
}
throw InputValidationError.fromCode('payload_json', ValidationErrorCodes.INVALID_JSON_IN_PAYLOAD_JSON);
}
async function readJsonBody(ctx: Context<HonoEnv>): Promise<unknown> {
if (isMultipart(ctx)) return readPayloadJson(ctx);
try {
const raw = await ctx.req.text();
return raw.trim().length === 0 ? {} : parseJsonPreservingLargeIntegers(raw);
} catch {
throw InputValidationError.fromCode('message_data', ValidationErrorCodes.INVALID_MESSAGE_DATA);
}
}
function parseWithSchema<T extends z.ZodType>(schema: T, value: unknown): z.output<T> {
const result = schema.safeParse(value);
if (!result.success) throw inputValidationErrorFromZodIssues(result.error.issues);
return result.data;
}
function SelfThreadMemberAlias(action: 'join' | 'leave'): MiddlewareHandler<HonoEnv> {
return async (ctx, next) => {
if (ctx.req.param('user_id') !== '@me') return next();
const {channel_id} = parseWithSchema(ChannelIdParam, {channel_id: ctx.req.param('channel_id')});
await ctx.get('threadService').members[action]({
viewer: viewerFromCtx(ctx),
user: ctx.get('user'),
channelId: createChannelID(channel_id),
});
return ctx.body(null, 204);
};
}
function normalizeForumEnvelope(value: unknown): unknown {
if (typeof value !== 'object' || value === null || Array.isArray(value)) return value;
const envelope = value as Record<string, unknown>;
return {...envelope, message: normalizeMessageRequestPayload(envelope.message)};
}
const ForumPostMultipartMessage = z.preprocess(
(value) =>
typeof value === 'object' && value !== null && !Array.isArray(value)
? ((value as Record<string, unknown>).message ?? {})
: {},
ForumThreadMessageRequest,
);
async function parseForumPostBody(ctx: Context<HonoEnv>, channelId: ChannelID): Promise<ForumPostInput> {
if (!isMultipart(ctx)) {
return parseWithSchema(StartForumThreadRequest, normalizeForumEnvelope(await readJsonBody(ctx))) as ForumPostInput;
}
let envelope: unknown = null;
const message = (await parseMultipartMessageData(
ctx,
ctx.get('user'),
channelId,
ForumPostMultipartMessage as unknown as z.ZodType<MessageRequest>,
{
onPayloadParsed(payload) {
envelope = payload;
},
},
)) as MessageRequest;
const fields = parseWithSchema(StartForumThreadRequest.omit({message: true}), envelope ?? {});
return {...fields, message};
}
export function ThreadController(app: HonoApp) {
app.post(
'/channels/:channel_id/messages/:message_id/threads',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREAD_CREATE),
LoginRequired,
Validator('param', ChannelIdMessageIdParam),
Validator('json', StartThreadFromMessageRequest),
OpenAPI({
operationId: 'start_thread_from_message',
summary: 'Start a thread from a message',
description:
'Creates a public thread from an existing message in a text channel. The thread shares the ID of the message, so a message can start one thread.',
responseSchema: ThreadChannelResponse,
statusCode: 201,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {channel_id, message_id} = ctx.req.valid('param');
const body = ctx.req.valid('json');
const thread = await ctx.get('threadService').creation.createFromMessage({
viewer: viewerFromCtx(ctx),
user: ctx.get('user'),
channelId: createChannelID(channel_id),
messageId: createMessageID(message_id),
input: {
name: body.name,
autoArchiveDuration: body.auto_archive_duration,
rateLimitPerUser: body.rate_limit_per_user,
},
requestCache: ctx.get('requestCache'),
auditLogReason: ctx.get('auditLogReason') ?? null,
});
return ctx.json(thread, 201);
},
);
app.post(
'/channels/:channel_id/threads',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREAD_CREATE),
LoginRequired,
Validator('param', ChannelIdParam),
OpenAPI({
operationId: 'start_thread',
summary: 'Start a thread',
description:
'Creates a thread that is not attached to an existing message. In a text channel the thread type is required. In a forum or media channel this creates a post, and the body carries the first message. The body can also be sent as multipart form data with the JSON in a payload_json field, and a post can attach files as files[n] parts.',
requestSchema: StartThreadRequestBody,
requestFormSchema: StartThreadMultipartRequest,
responseSchema: StartForumThreadResponse,
statusCode: 201,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const viewer = viewerFromCtx(ctx);
const user = ctx.get('user');
const channelId = createChannelID(ctx.req.valid('param').channel_id);
const creation = ctx.get('threadService').creation;
const parentAuth = await creation.authenticateParent(viewer, user.id, channelId);
const auditLogReason = ctx.get('auditLogReason') ?? null;
if (parentAuth.channel.isThreadOnly()) {
const body = await parseForumPostBody(ctx, channelId);
return ctx.json(
await creation.createForumPost({
viewer,
user,
parentAuth,
body,
requestCache: ctx.get('requestCache'),
auditLogReason,
}),
201,
);
}
if (!TEXT_THREAD_PARENT_CHANNEL_TYPES.has(parentAuth.channel.type)) throw new InvalidChannelTypeError();
const body = parseWithSchema(StartThreadRequest, await readJsonBody(ctx));
const thread = await creation.createTextThread({
user,
parentAuth,
type: body.type,
invitable: body.invitable,
input: {
name: body.name,
autoArchiveDuration: body.auto_archive_duration,
rateLimitPerUser: body.rate_limit_per_user,
},
auditLogReason,
});
return ctx.json(thread, 201);
},
);
app.get(
'/guilds/:guild_id/threads/active',
ChannelThreadsRouteGuard({botOnly: true}),
RateLimitMiddleware(RateLimitConfigs.GUILD_THREADS_ACTIVE),
LoginRequired,
BotOnly,
Validator('param', GuildIdParam),
OpenAPI({
operationId: 'list_guild_active_threads',
summary: 'List active guild threads',
description:
'Returns every active thread in the guild that the bot can view, newest first, with a thread member object for each thread the bot joined.',
responseSchema: ActiveThreadsResponse,
statusCode: 200,
security: ['botToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
return ctx.json(
await ctx.get('threadService').lists.listGuildActive({
userId: ctx.get('user').id,
guildId: createGuildID(ctx.req.valid('param').guild_id),
}),
);
},
);
app.get(
'/channels/:channel_id/threads/archived/public',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREADS_ARCHIVED_LIST),
LoginRequired,
Validator('param', ChannelIdParam),
Validator('query', ArchivedThreadsQuery),
OpenAPI({
operationId: 'list_public_archived_threads',
summary: 'List public archived threads',
description:
'Returns archived public threads of the channel, most recently archived first. Requires the read message history permission.',
responseSchema: ArchivedThreadsResponse,
statusCode: 200,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {before, limit} = ctx.req.valid('query');
return ctx.json(
await ctx.get('threadService').lists.listPublicArchived({
viewer: viewerFromCtx(ctx),
userId: ctx.get('user').id,
channelId: createChannelID(ctx.req.valid('param').channel_id),
before: before ? new Date(before) : undefined,
limit,
}),
);
},
);
app.get(
'/channels/:channel_id/threads/archived/private',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREADS_ARCHIVED_LIST),
LoginRequired,
Validator('param', ChannelIdParam),
Validator('query', ArchivedThreadsQuery),
OpenAPI({
operationId: 'list_private_archived_threads',
summary: 'List private archived threads',
description:
'Returns archived private threads of the text channel, most recently archived first. Requires the read message history and manage threads permissions.',
responseSchema: ArchivedThreadsResponse,
statusCode: 200,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {before, limit} = ctx.req.valid('query');
return ctx.json(
await ctx.get('threadService').lists.listPrivateArchived({
viewer: viewerFromCtx(ctx),
userId: ctx.get('user').id,
channelId: createChannelID(ctx.req.valid('param').channel_id),
before: before ? new Date(before) : undefined,
limit,
}),
);
},
);
app.get(
'/channels/:channel_id/users/@me/threads/archived/private',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREADS_ARCHIVED_LIST),
LoginRequired,
Validator('param', ChannelIdParam),
Validator('query', JoinedArchivedThreadsQuery),
OpenAPI({
operationId: 'list_joined_private_archived_threads',
summary: 'List joined private archived threads',
description:
'Returns archived private threads of the text channel that the current user joined, newest first. Requires the read message history permission.',
responseSchema: ArchivedThreadsResponse,
statusCode: 200,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {before, limit} = ctx.req.valid('query');
return ctx.json(
await ctx.get('threadService').lists.listJoinedPrivateArchived({
viewer: viewerFromCtx(ctx),
userId: ctx.get('user').id,
channelId: createChannelID(ctx.req.valid('param').channel_id),
before: before !== undefined ? createChannelID(before) : undefined,
limit,
}),
);
},
);
app.get(
'/channels/:channel_id/thread-members',
ChannelThreadsRouteGuard({botOnly: true}),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREAD_MEMBERS_LIST),
LoginRequired,
BotOnly,
Validator('param', ChannelIdParam),
Validator('query', ThreadMembersListQuery),
OpenAPI({
operationId: 'list_thread_members',
summary: 'List thread members',
description:
'Returns thread members ordered by user ID. Paginate with after and limit. Set with_member to include the guild member object of each thread member.',
responseSchema: ThreadMemberListResponse,
statusCode: 200,
security: ['botToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {with_member, after, limit} = ctx.req.valid('query');
return ctx.json(
await ctx.get('threadService').members.list({
viewer: viewerFromCtx(ctx),
user: ctx.get('user'),
channelId: createChannelID(ctx.req.valid('param').channel_id),
after: after !== undefined ? createUserID(after) : undefined,
limit,
withMember: with_member,
requestCache: ctx.get('requestCache'),
}),
);
},
);
app.put(
'/channels/:channel_id/thread-members/@me',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREAD_MEMBER_PUT),
LoginRequired,
Validator('param', ChannelIdParam),
Validator('query', ThreadLocationQuery),
OpenAPI({
operationId: 'join_thread',
summary: 'Join a thread',
description: 'Adds the current user to the thread. The thread must not be archived.',
responseSchema: null,
statusCode: 204,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
await ctx.get('threadService').members.join({
viewer: viewerFromCtx(ctx),
user: ctx.get('user'),
channelId: createChannelID(ctx.req.valid('param').channel_id),
});
return ctx.body(null, 204);
},
);
app.delete(
'/channels/:channel_id/thread-members/@me',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREAD_MEMBER_DELETE),
LoginRequired,
Validator('param', ChannelIdParam),
Validator('query', ThreadLocationQuery),
OpenAPI({
operationId: 'leave_thread',
summary: 'Leave a thread',
description: 'Removes the current user from the thread. The thread must not be archived.',
responseSchema: null,
statusCode: 204,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
await ctx.get('threadService').members.leave({
viewer: viewerFromCtx(ctx),
user: ctx.get('user'),
channelId: createChannelID(ctx.req.valid('param').channel_id),
});
return ctx.body(null, 204);
},
);
app.get(
'/channels/:channel_id/thread-members/:user_id',
ChannelThreadsRouteGuard({botOnly: true}),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREAD_MEMBER_GET),
LoginRequired,
BotOnly,
Validator('param', ChannelIdUserIdParam),
Validator('query', ThreadMemberGetQuery),
OpenAPI({
operationId: 'get_thread_member',
summary: 'Get a thread member',
description: 'Returns the thread member object of the user when the user is a member of the thread.',
responseSchema: ThreadMemberResponse,
statusCode: 200,
security: ['botToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {channel_id, user_id} = ctx.req.valid('param');
return ctx.json(
await ctx.get('threadService').members.get({
viewer: viewerFromCtx(ctx),
user: ctx.get('user'),
channelId: createChannelID(channel_id),
targetId: createUserID(user_id),
withMember: ctx.req.valid('query').with_member,
requestCache: ctx.get('requestCache'),
}),
);
},
);
app.put(
'/channels/:channel_id/thread-members/:user_id',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREAD_MEMBER_PUT),
LoginRequired,
SelfThreadMemberAlias('join'),
Validator('param', ChannelIdUserIdParam),
Validator('query', ThreadLocationQuery),
OpenAPI({
operationId: 'add_thread_member',
summary: 'Add a thread member',
description:
'Adds another guild member to the thread. Requires permission to send messages in threads, and the thread must not be archived.',
responseSchema: null,
statusCode: 204,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {channel_id, user_id} = ctx.req.valid('param');
await ctx.get('threadService').members.add({
viewer: viewerFromCtx(ctx),
user: ctx.get('user'),
channelId: createChannelID(channel_id),
targetId: createUserID(user_id),
});
return ctx.body(null, 204);
},
);
app.delete(
'/channels/:channel_id/thread-members/:user_id',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREAD_MEMBER_DELETE),
LoginRequired,
SelfThreadMemberAlias('leave'),
Validator('param', ChannelIdUserIdParam),
Validator('query', ThreadLocationQuery),
OpenAPI({
operationId: 'remove_thread_member',
summary: 'Remove a thread member',
description:
'Removes a member from the thread. Requires the manage threads permission, or being the creator of a private thread. The thread must not be archived.',
responseSchema: null,
statusCode: 204,
security: ['botToken', 'sessionToken'],
tags: TAGS,
experiment: EXPERIMENT,
}),
async (ctx) => {
const {channel_id, user_id} = ctx.req.valid('param');
await ctx.get('threadService').members.remove({
viewer: viewerFromCtx(ctx),
user: ctx.get('user'),
channelId: createChannelID(channel_id),
targetId: createUserID(user_id),
});
return ctx.body(null, 204);
},
);
}
@@ -0,0 +1,47 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {createChannelID} from '@app/api/BrandedTypes';
import {ChannelThreadsRouteGuard} from '@app/api/channel/threads/ChannelThreadsRouteGuard';
import {viewerFromCtx} from '@app/api/experiment/ChannelThreadsGate';
import {DefaultUserOnly, LoginRequired} from '@app/api/middleware/AuthMiddleware';
import {RateLimitMiddleware} from '@app/api/middleware/RateLimitMiddleware';
import {OpenAPI} from '@app/api/middleware/ResponseTypeMiddleware';
import {RateLimitConfigs} from '@app/api/RateLimitConfig';
import type {HonoApp} from '@app/api/types/HonoEnv';
import {Validator} from '@app/api/Validator';
import {ThreadMemberResponse} from '@fluxer/schema/src/domains/channel/ThreadSchemas';
import {ChannelIdParam} from '@fluxer/schema/src/domains/common/CommonParamSchemas';
import {ThreadMemberSettingsRequest} from '@fluxer/schema/src/domains/user/UserRequestSchemas';
export function ThreadMemberSettingsController(app: HonoApp) {
app.patch(
'/channels/:channel_id/thread-members/@me/settings',
ChannelThreadsRouteGuard(),
RateLimitMiddleware(RateLimitConfigs.CHANNEL_THREAD_MEMBER_SETTINGS),
LoginRequired,
DefaultUserOnly,
Validator('param', ChannelIdParam),
Validator('json', ThreadMemberSettingsRequest),
OpenAPI({
operationId: 'update_thread_member_settings',
summary: 'Update thread settings',
description:
"Updates the current user's notification settings for a thread they are a member of. Returns the thread member, or 204 when nothing changed.",
responseSchema: ThreadMemberResponse,
statusCode: [200, 204],
bodylessStatusCodes: [204],
security: ['sessionToken'],
tags: 'Channels',
experiment: 'channel_threads',
}),
async (ctx) => {
const member = await ctx.get('threadService').memberSettings.update({
viewer: viewerFromCtx(ctx),
userId: ctx.get('user').id,
channelId: createChannelID(ctx.req.valid('param').channel_id),
data: ctx.req.valid('json'),
});
return member ? ctx.json(member, 200) : ctx.body(null, 204);
},
);
}
@@ -3,9 +3,12 @@
import {CallController} from '@app/api/channel/controllers/CallController';
import {ChannelController} from '@app/api/channel/controllers/ChannelController';
import {ChannelFollowController} from '@app/api/channel/controllers/ChannelFollowController';
import {ForumController} from '@app/api/channel/controllers/ForumController';
import {MessageController} from '@app/api/channel/controllers/MessageController';
import {MessageInteractionController} from '@app/api/channel/controllers/MessageInteractionController';
import {StreamController} from '@app/api/channel/controllers/StreamController';
import {ThreadController} from '@app/api/channel/controllers/ThreadController';
import {ThreadMemberSettingsController} from '@app/api/channel/controllers/ThreadMemberSettingsController';
import type {HonoApp} from '@app/api/types/HonoEnv';
export function registerChannelControllers(app: HonoApp) {
@@ -15,4 +18,7 @@ export function registerChannelControllers(app: HonoApp) {
MessageController(app);
CallController(app);
StreamController(app);
ThreadController(app);
ForumController(app);
ThreadMemberSettingsController(app);
}
@@ -1,14 +1,15 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {ChannelID, GuildID, MessageID, UserID} from '@app/api/BrandedTypes';
import {type ChannelID, channelIdToMessageId, type GuildID, type MessageID, type UserID} from '@app/api/BrandedTypes';
import {
privateChannelFanOutTargets,
privateChannelLastMessageIdPatch,
privateChannelMetadataPatch,
} from '@app/api/channel/PrivateChannelSnapshot';
import {IChannelDataRepository} from '@app/api/channel/repositories/IChannelDataRepository';
import {type GuildChannelListMode, IChannelDataRepository} from '@app/api/channel/repositories/IChannelDataRepository';
import {
BatchBuilder,
executeConditional,
fetchMany,
fetchManyInChunks,
fetchOne,
@@ -18,10 +19,13 @@ import {Db} from '@app/api/database/CassandraTypes';
import {buildPatchFromData, executeVersionedUpdate} from '@app/api/database/CassandraVersionedUpdate';
import type {ChannelRow} from '@app/api/database/types/ChannelTypes';
import {CHANNEL_COLUMNS} from '@app/api/database/types/ChannelTypes';
import type {ThreadStatsRow} from '@app/api/database/types/ThreadTypes';
import {guildActive, isTainted} from '@app/api/experiment/ChannelThreadsGate';
import {Logger} from '@app/api/Logger';
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
import {Channel} from '@app/api/models/Channel';
import {Channels, ChannelsByGuild, PrivateChannels} from '@app/api/Tables';
import {Channels, ChannelsByGuild, PrivateChannels, ThreadOnlyChannelsByGuild, ThreadStats} from '@app/api/Tables';
import {THREAD_CHANNEL_TYPES, THREAD_ONLY_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
const FETCH_CHANNEL_BY_ID = Channels.select({
where: [Channels.where.eq('channel_id'), Channels.where.eq('soft_deleted')],
@@ -33,6 +37,11 @@ const FETCH_CHANNELS_BY_IDS = Channels.select({
const FETCH_GUILD_CHANNELS_BY_GUILD_ID = ChannelsByGuild.select({
where: ChannelsByGuild.where.eq('guild_id'),
});
const FETCH_THREAD_ONLY_CHANNELS_BY_GUILD_ID = ThreadOnlyChannelsByGuild.select({
where: ThreadOnlyChannelsByGuild.where.eq('guild_id'),
});
const THREAD_STATS_CAS_ATTEMPTS = 16;
const FETCH_THREAD_STATS = ThreadStats.select({where: ThreadStats.where.eq('thread_id'), limit: 1});
const FETCH_OPEN_PRIVATE_CHANNEL_TARGET = PrivateChannels.selectCql({
columns: ['user_id'],
where: [PrivateChannels.where.eq('user_id'), PrivateChannels.where.eq('channel_id')],
@@ -70,12 +79,14 @@ export class ChannelDataRepository extends IChannelDataRepository {
Channels,
{initialData: oldData},
);
if (data.guild_id) {
if (data.guild_id && !THREAD_CHANNEL_TYPES.has(data.type)) {
await upsertOne(
ChannelsByGuild.upsertAll({
guild_id: data.guild_id,
channel_id: channelId,
}),
THREAD_ONLY_CHANNEL_TYPES.has(data.type)
? ThreadOnlyChannelsByGuild.upsertAll({guild_id: data.guild_id, channel_id: channelId})
: ChannelsByGuild.upsertAll({
guild_id: data.guild_id,
channel_id: channelId,
}),
);
}
const finalRow: ChannelRow = {...data, version: result.finalVersion ?? 0};
@@ -83,7 +94,7 @@ export class ChannelDataRepository extends IChannelDataRepository {
return new Channel(finalRow);
}
async updateLastMessageId(channelId: ChannelID, messageId: MessageID): Promise<void> {
async updateLastMessageId(channelId: ChannelID, messageId: MessageID, opts?: {isInsert?: boolean}): Promise<void> {
this.requestCache?.channels.delete(channelId);
const existing = await fetchOne<ChannelRow>(
FETCH_CHANNEL_BY_ID.bind({
@@ -92,6 +103,9 @@ export class ChannelDataRepository extends IChannelDataRepository {
}),
);
if (!existing) return;
if (opts?.isInsert && THREAD_CHANNEL_TYPES.has(existing.type) && messageId !== channelIdToMessageId(channelId)) {
await this.adjustThreadStats(channelId, 1, 1);
}
const prev = existing.last_message_id ?? null;
if (prev !== null && messageId <= prev) return;
await upsertOne(
@@ -100,6 +114,34 @@ export class ChannelDataRepository extends IChannelDataRepository {
void this.fanOutPrivateChannelLastMessageId(existing, messageId);
}
async adjustThreadStats(threadId: ChannelID, messageDelta: number, sentDelta: number): Promise<void> {
for (let attempt = 0; attempt < THREAD_STATS_CAS_ATTEMPTS; attempt++) {
const stats = await fetchOne<ThreadStatsRow>(FETCH_THREAD_STATS.bind({thread_id: threadId}));
const messageCount = Math.max(0, (stats?.message_count ?? 0) + messageDelta);
const totalMessageSent = Math.max(0, (stats?.total_message_sent ?? 0) + sentDelta);
const applied = await executeConditional(
stats
? ThreadStats.conditionalPatchByPk(
{thread_id: threadId},
{message_count: Db.set(messageCount), total_message_sent: Db.set(totalMessageSent)},
{message_count: stats.message_count ?? null, total_message_sent: stats.total_message_sent ?? null},
)
: ThreadStats.insertIfNotExists({
thread_id: threadId,
message_count: messageCount,
total_message_sent: totalMessageSent,
}),
);
if (applied) return;
}
Logger.warn({threadId: threadId.toString()}, 'Gave up adjusting thread stats under contention');
}
async patchIndexedAt(channelId: ChannelID, indexedAt: Date): Promise<void> {
this.requestCache?.channels.delete(channelId);
await upsertOne(Channels.patchByPk({channel_id: channelId, soft_deleted: false}, {indexed_at: Db.set(indexedAt)}));
}
private async writeThroughPrivateChannelMetadata(row: ChannelRow): Promise<void> {
try {
const targets = await this.listOpenPrivateChannelTargets(row);
@@ -175,7 +217,7 @@ export class ChannelDataRepository extends IChannelDataRepository {
);
}
async delete(channelId: ChannelID, guildId?: GuildID): Promise<void> {
async delete(channelId: ChannelID, guildId?: GuildID, type?: number): Promise<void> {
this.requestCache?.channels.delete(channelId);
const batch = new BatchBuilder();
batch.addPrepared(
@@ -184,23 +226,30 @@ export class ChannelDataRepository extends IChannelDataRepository {
soft_deleted: false,
}),
);
if (guildId) {
if (guildId && (type === undefined || !THREAD_CHANNEL_TYPES.has(type))) {
batch.addPrepared(
ChannelsByGuild.deleteByPk({
guild_id: guildId,
channel_id: channelId,
}),
type !== undefined && THREAD_ONLY_CHANNEL_TYPES.has(type)
? ThreadOnlyChannelsByGuild.deleteByPk({guild_id: guildId, channel_id: channelId})
: ChannelsByGuild.deleteByPk({
guild_id: guildId,
channel_id: channelId,
}),
);
}
await batch.execute();
}
async listGuildChannels(guildId: GuildID): Promise<Array<Channel>> {
const guildChannels = await fetchMany<{
channel_id: bigint;
}>(FETCH_GUILD_CHANNELS_BY_GUILD_ID.bind({guild_id: guildId}));
if (guildChannels.length === 0) return [];
const channelIds = guildChannels.map((c) => c.channel_id);
async listGuildChannels(guildId: GuildID, mode: GuildChannelListMode): Promise<Array<Channel>> {
const includeThreadOnly =
mode === 'enrolled' ? guildActive(guildId) : await isTainted(guildId, {fresh: mode === 'complete'});
const [guildChannels, threadOnlyChannels] = await Promise.all([
fetchMany<{channel_id: bigint}>(FETCH_GUILD_CHANNELS_BY_GUILD_ID.bind({guild_id: guildId})),
includeThreadOnly
? fetchMany<{channel_id: bigint}>(FETCH_THREAD_ONLY_CHANNELS_BY_GUILD_ID.bind({guild_id: guildId}))
: Promise.resolve([]),
]);
if (guildChannels.length === 0 && threadOnlyChannels.length === 0) return [];
const channelIds = [...guildChannels, ...threadOnlyChannels].map((c) => c.channel_id);
const channels = await fetchManyInChunks<ChannelRow>(FETCH_CHANNELS_BY_IDS, channelIds, (chunk) => ({
channel_ids: chunk,
soft_deleted: false,
@@ -218,9 +267,12 @@ export class ChannelDataRepository extends IChannelDataRepository {
}
async countGuildChannels(guildId: GuildID): Promise<number> {
const guildChannels = await fetchMany<{
channel_id: bigint;
}>(FETCH_GUILD_CHANNELS_BY_GUILD_ID.bind({guild_id: guildId}));
return guildChannels.length;
const [guildChannels, threadOnlyChannels] = await Promise.all([
fetchMany<{channel_id: bigint}>(FETCH_GUILD_CHANNELS_BY_GUILD_ID.bind({guild_id: guildId})),
guildActive(guildId)
? fetchMany<{channel_id: bigint}>(FETCH_THREAD_ONLY_CHANNELS_BY_GUILD_ID.bind({guild_id: guildId}))
: Promise.resolve([]),
]);
return guildChannels.length + threadOnlyChannels.length;
}
}
@@ -5,12 +5,15 @@ import {CrosspostedMessageRepository} from '@app/api/channel/repositories/Crossp
import {IChannelRepositoryAggregate} from '@app/api/channel/repositories/IChannelRepositoryAggregate';
import {MessageInteractionRepository} from '@app/api/channel/repositories/MessageInteractionRepository';
import {MessageRepository} from '@app/api/channel/repositories/MessageRepository';
import {ThreadRepository} from '@app/api/channel/repositories/ThreadRepository';
import {enqueueRepairThreadIndexes} from '@app/api/channel/threads/ThreadJobs';
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
export class ChannelRepository extends IChannelRepositoryAggregate {
readonly channelData: ChannelDataRepository;
readonly messages: MessageRepository;
readonly messageInteractions: MessageInteractionRepository;
readonly threads: ThreadRepository;
readonly crossposts: CrosspostedMessageRepository;
constructor(requestCache?: RequestCache) {
@@ -18,6 +21,7 @@ export class ChannelRepository extends IChannelRepositoryAggregate {
this.channelData = new ChannelDataRepository(requestCache);
this.messages = new MessageRepository(this.channelData);
this.messageInteractions = new MessageInteractionRepository(this.messages);
this.threads = new ThreadRepository(this.channelData, this.messages, enqueueRepairThreadIndexes);
this.crossposts = new CrosspostedMessageRepository();
}
}
@@ -4,18 +4,22 @@ import type {ChannelID, GuildID, MessageID} from '@app/api/BrandedTypes';
import type {ChannelRow} from '@app/api/database/types/ChannelTypes';
import type {Channel} from '@app/api/models/Channel';
export type GuildChannelListMode = 'enrolled' | 'maintenance' | 'complete';
export abstract class IChannelDataRepository {
abstract findUnique(channelId: ChannelID): Promise<Channel | null>;
abstract upsert(data: ChannelRow): Promise<Channel>;
abstract upsert(data: ChannelRow, oldData?: ChannelRow | null): Promise<Channel>;
abstract updateLastMessageId(channelId: ChannelID, messageId: MessageID): Promise<void>;
abstract updateLastMessageId(channelId: ChannelID, messageId: MessageID, opts?: {isInsert?: boolean}): Promise<void>;
abstract delete(channelId: ChannelID, guildId?: GuildID): Promise<void>;
abstract delete(channelId: ChannelID, guildId?: GuildID, type?: number): Promise<void>;
abstract listGuildChannels(guildId: GuildID): Promise<Array<Channel>>;
abstract listGuildChannels(guildId: GuildID, mode: GuildChannelListMode): Promise<Array<Channel>>;
abstract listChannels(channelIds: Array<ChannelID>): Promise<Array<Channel>>;
abstract countGuildChannels(guildId: GuildID): Promise<number>;
abstract patchIndexedAt(channelId: ChannelID, indexedAt: Date): Promise<void>;
}
@@ -4,10 +4,12 @@ import type {IChannelDataRepository} from '@app/api/channel/repositories/IChanne
import type {ICrosspostedMessageRepository} from '@app/api/channel/repositories/ICrosspostedMessageRepository';
import type {IMessageInteractionRepository} from '@app/api/channel/repositories/IMessageInteractionRepository';
import type {IMessageRepository} from '@app/api/channel/repositories/IMessageRepository';
import type {IThreadRepository} from '@app/api/channel/repositories/IThreadRepository';
export abstract class IChannelRepositoryAggregate {
abstract readonly channelData: IChannelDataRepository;
abstract readonly messages: IMessageRepository;
abstract readonly messageInteractions: IMessageInteractionRepository;
abstract readonly threads: IThreadRepository;
abstract readonly crossposts: ICrosspostedMessageRepository;
}
@@ -9,6 +9,11 @@ export interface ListMessagesOptions {
immediateAfter?: boolean;
}
export interface UpsertMessageOptions {
isInsert?: boolean;
skipParentLastMessageId?: boolean;
}
export abstract class IMessageRepository {
abstract listMessages(
channelId: ChannelID,
@@ -20,7 +25,7 @@ export abstract class IMessageRepository {
abstract getMessage(channelId: ChannelID, messageId: MessageID): Promise<Message | null>;
abstract upsertMessage(data: MessageRow, oldData?: MessageRow | null): Promise<Message>;
abstract upsertMessage(data: MessageRow, oldData?: MessageRow | null, opts?: UpsertMessageOptions): Promise<Message>;
abstract updateEmbeds(message: Message): Promise<void>;
@@ -0,0 +1,172 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {ChannelID, GuildID, MessageID, UserID} from '@app/api/BrandedTypes';
import type {ChannelRow} from '@app/api/database/types/ChannelTypes';
import type {GuildThreadStateRow, ThreadParentConfigRow, ThreadStateRow} from '@app/api/database/types/ThreadTypes';
import type {MuteConfig} from '@app/api/database/types/UserTypes';
import type {ThreadMember} from '@app/api/models/ThreadMember';
import type {ThreadParentConfig} from '@app/api/models/ThreadParentConfig';
import type {ThreadState} from '@app/api/models/ThreadState';
import type {ThreadStats} from '@app/api/models/ThreadStats';
export interface CreateThreadMember {
userId: UserID;
flags: number;
}
export interface CreateThreadParams {
channel: ChannelRow;
parentType: number;
autoArchiveDuration: number;
invitable: boolean | null;
flags: number;
appliedTags: Array<bigint>;
hasStarter: boolean;
createdAt: Date;
members: Array<CreateThreadMember>;
}
export type ThreadStatePatch = Partial<
Pick<
ThreadStateRow,
'archived' | 'locked' | 'invitable' | 'auto_archive_duration' | 'archive_timestamp' | 'flags' | 'applied_tags'
>
>;
export interface ThreadStateTransition {
previous: ThreadState;
state: ThreadState;
}
export interface ArchivedThreadPage {
threads: Array<ThreadState>;
hasMore: boolean;
}
export interface ThreadMemberSettingsPatch {
flags?: number;
muted?: boolean;
muteConfig?: MuteConfig | null;
}
export interface ThreadMemberAddResult {
added: Array<ThreadMember>;
state: ThreadState;
}
export interface ThreadMemberRemoveResult {
removed: Array<ThreadMember>;
state: ThreadState | null;
}
export type ThreadParentConfigPatch = Partial<Omit<ThreadParentConfigRow, 'guild_id' | 'channel_id'>>;
export abstract class IThreadRepository {
abstract getState(threadId: ChannelID): Promise<ThreadState | null>;
abstract getStates(threadIds: Array<ChannelID>): Promise<Array<ThreadState>>;
abstract getStats(threadId: ChannelID): Promise<ThreadStats>;
abstract getStatsMany(threadIds: Array<ChannelID>): Promise<Map<ChannelID, ThreadStats>>;
abstract adjustMessageCount(threadId: ChannelID, delta: number): Promise<void>;
abstract create(params: CreateThreadParams): Promise<ThreadState>;
abstract updateState(
threadId: ChannelID,
mutate: (current: ThreadState) => ThreadStatePatch | null,
): Promise<ThreadStateTransition | null>;
abstract claimForumPin(parentId: ChannelID, threadId: ChannelID): Promise<boolean>;
abstract releaseForumPin(parentId: ChannelID, threadId: ChannelID): Promise<void>;
abstract getForumPin(parentId: ChannelID): Promise<ChannelID | null>;
abstract countActiveThreads(guildId: GuildID): Promise<number>;
abstract listActiveThreads(guildId: GuildID): Promise<Array<ThreadState>>;
abstract listArchivedThreads(
parentId: ChannelID,
isPrivate: boolean,
opts: {before?: Date; limit: number},
): Promise<ArchivedThreadPage>;
abstract listJoinedPrivateArchivedThreads(
userId: UserID,
guildId: GuildID,
parentId: ChannelID,
opts: {before?: ChannelID; limit: number},
): Promise<ArchivedThreadPage>;
abstract listJoinedThreadIds(userId: UserID, guildId: GuildID): Promise<Array<ChannelID>>;
abstract listJoinedPrivateThreadIds(
userId: UserID,
guildId: GuildID,
parentId: ChannelID,
limit: number,
): Promise<Array<ChannelID>>;
abstract repairThreadIndexes(threadIds: Array<ChannelID>): Promise<void>;
abstract getMember(threadId: ChannelID, userId: UserID): Promise<ThreadMember | null>;
abstract getMembers(threadId: ChannelID, userIds: Array<UserID>): Promise<Array<ThreadMember>>;
abstract listMembers(threadId: ChannelID, opts: {after?: UserID; limit: number}): Promise<Array<ThreadMember>>;
abstract addMembers(
threadId: ChannelID,
members: Array<CreateThreadMember>,
opts?: {joinTimestamp?: Date},
): Promise<ThreadMemberAddResult | null>;
abstract removeMembers(threadId: ChannelID, userIds: Array<UserID>): Promise<ThreadMemberRemoveResult>;
abstract updateMemberSettings(expected: ThreadMember, patch: ThreadMemberSettingsPatch): Promise<ThreadMember | null>;
abstract listThreadIdsByParent(
parentId: ChannelID,
opts: {after?: ChannelID; limit: number},
): Promise<Array<ChannelID>>;
abstract listParentThreads(parentId: ChannelID): Promise<Array<{threadId: ChannelID; type: number}>>;
abstract setThreadType(threadId: ChannelID, type: number): Promise<ThreadState | null>;
abstract listGuildThreadIds(
guildId: GuildID,
opts?: {activeSince?: Date; parents?: ReadonlyArray<{id: ChannelID; type: number}>},
): Promise<Array<ChannelID>>;
abstract purgeThread(threadId: ChannelID): Promise<void>;
abstract revertParentLastMessageId(
parentId: ChannelID,
threadId: ChannelID,
previous: MessageID | null,
): Promise<void>;
abstract purgeGuild(guildId: GuildID): Promise<void>;
abstract getGuildMarker(guildId: GuildID): Promise<GuildThreadStateRow | null>;
abstract ensureGuildMarker(guildId: GuildID, opts?: {permsSeededAt?: Date}): Promise<void>;
abstract markGuildPermsSeeded(guildId: GuildID, at: Date): Promise<void>;
abstract markGuildSearchBackfilled(guildId: GuildID, at: Date): Promise<void>;
abstract clearGuildSearchBackfilled(guildId: GuildID): Promise<void>;
abstract getParentConfig(guildId: GuildID, channelId: ChannelID): Promise<ThreadParentConfig | null>;
abstract listParentConfigs(guildId: GuildID): Promise<Array<ThreadParentConfig>>;
abstract patchParentConfig(guildId: GuildID, channelId: ChannelID, patch: ThreadParentConfigPatch): Promise<void>;
abstract deleteParentConfig(guildId: GuildID, channelId: ChannelID): Promise<void>;
}
@@ -2,7 +2,11 @@
import type {AttachmentID, ChannelID, MessageID, UserID} from '@app/api/BrandedTypes';
import type {ChannelDataRepository} from '@app/api/channel/repositories/ChannelDataRepository';
import {IMessageRepository, type ListMessagesOptions} from '@app/api/channel/repositories/IMessageRepository';
import {
IMessageRepository,
type ListMessagesOptions,
type UpsertMessageOptions,
} from '@app/api/channel/repositories/IMessageRepository';
import {MessageAttachmentRepository} from '@app/api/channel/repositories/message/MessageAttachmentRepository';
import {MessageAuthorRepository} from '@app/api/channel/repositories/message/MessageAuthorRepository';
import {MessageDataRepository} from '@app/api/channel/repositories/message/MessageDataRepository';
@@ -40,10 +44,14 @@ export class MessageRepository extends IMessageRepository {
return this.dataRepo.getMessage(channelId, messageId);
}
async upsertMessage(data: MessageRow, oldData?: MessageRow | null): Promise<Message> {
async upsertMessage(data: MessageRow, oldData?: MessageRow | null, opts?: UpsertMessageOptions): Promise<Message> {
const message = await this.dataRepo.upsertMessage(data, oldData);
if (!oldData) {
await this.channelDataRepo.updateLastMessageId(data.channel_id, data.message_id);
if (!oldData && !opts?.skipParentLastMessageId) {
await this.channelDataRepo.updateLastMessageId(
data.channel_id,
data.message_id,
opts?.isInsert ? {isInsert: true} : undefined,
);
}
return message;
}
@@ -0,0 +1,992 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {spawnSync} from 'node:child_process';
import {readFileSync} from 'node:fs';
import {createServer} from 'node:net';
import {fileURLToPath} from 'node:url';
import {
type ChannelID,
channelIdToMessageId,
createChannelID,
createGuildID,
createMessageID,
createUserID,
type GuildID,
} from '@app/api/BrandedTypes';
import {ChannelRepository} from '@app/api/channel/repositories/ChannelRepository';
import type {CreateThreadParams} from '@app/api/channel/repositories/IThreadRepository';
import {ThreadRepository} from '@app/api/channel/repositories/ThreadRepository';
import {
type CassandraQueryExecutorForTesting,
deleteOneOrMany,
executeConditional,
fetchMany,
fetchOne,
setCassandraQueryExecutorForTesting,
upsertOne,
} from '@app/api/database/CassandraQueryExecution';
import type {CassandraParams, KvQueryMeta, PreparedQuery} from '@app/api/database/CassandraTypes';
import {ensurePostgresKvSchema, PostgresKvQueryExecutor} from '@app/api/database/PostgresKvQueryExecutor';
import {CHANNEL_COLUMNS, type ChannelRow} from '@app/api/database/types/ChannelTypes';
import {MESSAGE_COLUMNS} from '@app/api/database/types/MessageTypes';
import {
clearChannelThreadsTaintCacheForTesting,
syncChannelThreadsConfig,
} from '@app/api/experiment/ChannelThreadsGate';
import {
ActiveThreadsByGuild,
ArchivedThreadsByParent,
Channels,
ChannelsByGuild,
ForumPinnedThread,
GuildThreadState,
ThreadMembers,
ThreadMembersByUser,
ThreadOnlyChannelsByGuild,
ThreadParentConfig,
ThreadState,
ThreadStats,
ThreadsByParent,
} from '@app/api/Tables';
import {startDockerContainer} from '@app/api/test/DockerTestContainer';
import {InMemoryCassandraQueryExecutor} from '@app/api/test/InMemoryCassandraQueryExecutor';
import {ChannelTypes} from '@fluxer/constants/src/ChannelConstants';
import {ChannelFlags, MAX_THREAD_MEMBERS} from '@fluxer/constants/src/ThreadConstants';
import {MaxThreadMembersError} from '@fluxer/errors/src/domains/channel/MaxThreadMembersError';
import {ThreadAlreadyCreatedForMessageError} from '@fluxer/errors/src/domains/channel/ThreadAlreadyCreatedForMessageError';
import {
type ChannelThreadsConfig,
ChannelThreadsConfigSchema,
} from '@fluxer/schema/src/domains/admin/ChannelThreadsSchemas';
import {createSnowflakeFromTimestamp} from '@fluxer/snowflake/src/Snowflake';
import {getDefaultPostgresClient, initPostgres, shutdownPostgres} from '@pkgs/postgres/src/Client';
import {afterAll, afterEach, beforeAll, beforeEach, describe, expect, it, vi} from 'vitest';
const GUILD_ID = createGuildID(1_900_000_000_000_000_000n);
const PARENT_ID = createChannelID(1_900_000_000_000_000_100n);
const FORUM_ID = createChannelID(1_900_000_000_000_000_200n);
const OWNER_ID = createUserID(1_900_000_000_000_001_000n);
class RecordingExecutor implements CassandraQueryExecutorForTesting {
readonly statements: Array<string> = [];
constructor(private readonly inner: CassandraQueryExecutorForTesting) {}
async executeQuery<T = Record<string, unknown>, P extends CassandraParams = CassandraParams>(
query: PreparedQuery<P>,
): Promise<Array<T>> {
const meta = query.kvMeta;
this.record(meta, meta?.conditions || meta?.ifNotExists || meta?.batchEntries ? 'cas' : undefined);
return this.inner.executeQuery<T, P>(query);
}
async executeBatch(
queries: Array<{query: string; params: object; meta?: KvQueryMeta}>,
atomic?: boolean,
): Promise<void> {
for (const entry of queries) this.record(entry.meta);
await this.inner.executeBatch(queries, atomic);
}
count(statement: string): number {
return this.statements.filter((entry) => entry === statement).length;
}
private record(meta: KvQueryMeta | undefined, prefix?: string): void {
if (!meta) return;
this.statements.push(`${prefix ?? meta.action}:${meta.table.name}`);
}
}
function parseConfig(raw: string | null): ChannelThreadsConfig {
return ChannelThreadsConfigSchema.parse(raw ? JSON.parse(raw) : {});
}
function setThreadsConfig(patch: Partial<ChannelThreadsConfig> | null): void {
clearChannelThreadsTaintCacheForTesting();
syncChannelThreadsConfig(patch === null ? null : JSON.stringify(patch), parseConfig);
}
function channelRow(channelId: ChannelID, type: number, overrides: Partial<ChannelRow> = {}): ChannelRow {
return {
channel_id: channelId,
guild_id: GUILD_ID,
type,
name: `channel-${channelId}`,
topic: null,
icon_hash: null,
url: null,
parent_id: null,
position: 0,
owner_id: null,
recipient_ids: null,
nsfw: false,
content_warning_level: null,
content_warning_text: null,
rate_limit_per_user: 0,
bitrate: null,
user_limit: null,
voice_connection_limit: null,
rtc_region: null,
last_message_id: null,
last_pin_timestamp: null,
permission_overwrites: null,
nicks: null,
soft_deleted: false,
indexed_at: null,
version: 1,
...overrides,
};
}
let lastId = 0n;
function freshThreadId(offsetMs = 0): ChannelID {
const base = createSnowflakeFromTimestamp(Date.now() + offsetMs);
if (offsetMs !== 0) return createChannelID(base);
lastId = base > lastId ? base : lastId + 1n;
return createChannelID(lastId);
}
function createParams(threadId: ChannelID, overrides: Partial<CreateThreadParams> = {}): CreateThreadParams {
const type = overrides.channel?.type ?? ChannelTypes.PUBLIC_THREAD;
return {
channel: channelRow(threadId, type, {parent_id: PARENT_ID, owner_id: OWNER_ID, indexed_at: new Date()}),
parentType: ChannelTypes.GUILD_TEXT,
autoArchiveDuration: 4320,
invitable: null,
flags: 0,
appliedTags: [],
hasStarter: false,
createdAt: new Date(),
members: [{userId: OWNER_ID, flags: 1}],
...overrides,
};
}
function describeThreadRepository(backend: string, makeExecutor: () => Promise<CassandraQueryExecutorForTesting>) {
describe(`ThreadRepository (${backend})`, () => {
let executor: RecordingExecutor;
let repositories: ChannelRepository;
beforeEach(async () => {
executor = new RecordingExecutor(await makeExecutor());
setCassandraQueryExecutorForTesting(executor);
setThreadsConfig(null);
repositories = new ChannelRepository();
await upsertOne(Channels.upsertAll(channelRow(PARENT_ID, ChannelTypes.GUILD_TEXT)));
await upsertOne(ChannelsByGuild.upsertAll({guild_id: GUILD_ID, channel_id: PARENT_ID}));
executor.statements.length = 0;
});
afterEach(() => {
setThreadsConfig(null);
});
async function taint(guildId: GuildID = GUILD_ID): Promise<void> {
setThreadsConfig({enabled: true, ever_enabled: true, enabled_guild_ids: [guildId.toString()]});
await upsertOne(
GuildThreadState.upsertAll({
guild_id: guildId,
first_active_at: new Date(),
perms_seeded_at: null,
search_backfilled_at: null,
}),
);
}
it('writes a thread into its side tables and never into the guild channel index', async () => {
const threadId = freshThreadId();
const state = await repositories.threads.create(createParams(threadId));
expect(state.stateVersion).toBe(1);
expect(state.memberCount).toBe(1);
expect(state.memberIdsPreview).toEqual([OWNER_ID]);
const channel = await repositories.channelData.findUnique(threadId);
expect(channel?.isThread()).toBe(true);
expect(await repositories.channelData.listGuildChannels(GUILD_ID, 'enrolled')).toHaveLength(1);
expect(
await fetchOne(
ChannelsByGuild.selectCql({
where: [ChannelsByGuild.where.eq('guild_id'), ChannelsByGuild.where.eq('channel_id')],
}),
{guild_id: GUILD_ID, channel_id: threadId},
),
).toBeNull();
expect((await repositories.threads.listActiveThreads(GUILD_ID)).map((t) => t.threadId)).toEqual([threadId]);
expect(await repositories.threads.listThreadIdsByParent(PARENT_ID, {limit: 10})).toEqual([threadId]);
expect((await repositories.threads.getParentConfig(GUILD_ID, PARENT_ID))?.hasThreads).toBe(true);
expect((await repositories.threads.getMember(threadId, OWNER_ID))?.flags).toBe(1);
});
it('lets exactly one of two concurrent creates win', async () => {
const threadId = freshThreadId();
const results = await Promise.allSettled([
repositories.threads.create(createParams(threadId)),
repositories.threads.create(createParams(threadId)),
]);
expect(results.filter((result) => result.status === 'fulfilled')).toHaveLength(1);
const rejected = results.find((result) => result.status === 'rejected') as PromiseRejectedResult;
expect(rejected.reason).toBeInstanceOf(ThreadAlreadyCreatedForMessageError);
});
it('refuses a create over a fresh dangling state and repairs one older than 30 seconds', async () => {
const threadId = freshThreadId();
const dangling = {
thread_id: threadId,
guild_id: GUILD_ID,
parent_id: PARENT_ID,
type: ChannelTypes.PUBLIC_THREAD,
archived: false,
locked: false,
invitable: null,
auto_archive_duration: 4320,
archive_timestamp: new Date(),
created_at: new Date(),
flags: 0,
applied_tags: null,
member_count: 0,
member_ids_preview: null,
has_starter: false,
state_version: 1,
};
expect(await executeConditional(ThreadState.insertIfNotExists(dangling))).toBe(true);
await expect(repositories.threads.create(createParams(threadId))).rejects.toBeInstanceOf(
ThreadAlreadyCreatedForMessageError,
);
const staleId = freshThreadId();
expect(
await executeConditional(
ThreadState.insertIfNotExists({...dangling, thread_id: staleId, created_at: new Date(Date.now() - 31_000)}),
),
).toBe(true);
const repaired = await repositories.threads.create(createParams(staleId));
expect(repaired.memberCount).toBe(1);
expect(await repositories.channelData.findUnique(staleId)).not.toBeNull();
});
it('clears the crashed creator memberships when repairing a dangling create', async () => {
const threadId = freshThreadId();
const crashed = createUserID(1_900_000_000_000_003_000n);
const createdAt = new Date(Date.now() - 31_000);
expect(
await executeConditional(
ThreadState.insertIfNotExists({
thread_id: threadId,
guild_id: GUILD_ID,
parent_id: PARENT_ID,
type: ChannelTypes.PRIVATE_THREAD,
archived: false,
locked: false,
invitable: true,
auto_archive_duration: 4320,
archive_timestamp: createdAt,
created_at: createdAt,
flags: 0,
applied_tags: null,
member_count: 1,
member_ids_preview: [crashed],
has_starter: false,
state_version: 1,
}),
),
).toBe(true);
const memberRow = {
thread_id: threadId,
guild_id: GUILD_ID,
parent_id: PARENT_ID,
join_timestamp: createdAt,
flags: 1,
muted: false,
mute_config: null,
};
await upsertOne(ThreadMembers.upsertAll({...memberRow, user_id: crashed}));
await upsertOne(ThreadMembers.upsertAll({...memberRow, user_id: OWNER_ID}));
await upsertOne(
ThreadMembersByUser.upsertAll({
user_id: crashed,
guild_id: GUILD_ID,
parent_id: PARENT_ID,
is_private: true,
thread_id: threadId,
}),
);
const repaired = await repositories.threads.create(createParams(threadId));
expect(repaired.memberCount).toBe(1);
const members = await repositories.threads.listMembers(threadId, {limit: 10});
expect(members.map((member) => member.userId)).toEqual([OWNER_ID]);
expect(members[0]!.joinTimestamp.getTime()).not.toBe(createdAt.getTime());
expect(await repositories.threads.listJoinedThreadIds(crashed, GUILD_ID)).toEqual([]);
});
it('removes a private membership index row when the thread state is gone', async () => {
const threadId = freshThreadId();
const params = createParams(threadId);
params.channel.type = ChannelTypes.PRIVATE_THREAD;
await repositories.threads.create(params);
await deleteOneOrMany(ThreadState.deleteByPk({thread_id: threadId}));
const result = await repositories.threads.removeMembers(threadId, [OWNER_ID]);
expect(result.removed.map((member) => member.userId)).toEqual([OWNER_ID]);
expect(await repositories.threads.listJoinedThreadIds(OWNER_ID, GUILD_ID)).toEqual([]);
});
it('ignores a dangling threads_by_parent row in lists but enumerates it for maintenance', async () => {
await taint();
const threadId = freshThreadId();
await upsertOne(
ThreadsByParent.upsertAll({
parent_id: PARENT_ID,
thread_id: threadId,
guild_id: GUILD_ID,
type: ChannelTypes.PUBLIC_THREAD,
}),
);
expect(await repositories.threads.listActiveThreads(GUILD_ID)).toEqual([]);
expect(await repositories.threads.listGuildThreadIds(GUILD_ID)).toEqual([threadId]);
});
it('moves index rows on archive and unarchive and clears the pin', async () => {
const threadId = freshThreadId();
await repositories.threads.create(createParams(threadId, {flags: ChannelFlags.PINNED}));
expect(await repositories.threads.claimForumPin(FORUM_ID, threadId)).toBe(true);
await upsertOne(ThreadState.patchByPk({thread_id: threadId}, {parent_id: {kind: 'set', value: FORUM_ID}}));
const archived = await repositories.threads.updateState(threadId, () => ({archived: true}));
expect(archived?.state.archived).toBe(true);
expect(archived?.state.isPinned).toBe(false);
expect(archived?.state.stateVersion).toBe(2);
expect(await repositories.threads.getForumPin(FORUM_ID)).toBeNull();
expect(await repositories.threads.listActiveThreads(GUILD_ID)).toEqual([]);
const page = await repositories.threads.listArchivedThreads(FORUM_ID, false, {limit: 10});
expect(page.threads.map((t) => t.threadId)).toEqual([threadId]);
const unarchived = await repositories.threads.updateState(threadId, () => ({archived: false}));
expect(unarchived?.state.archived).toBe(false);
expect((await repositories.threads.listActiveThreads(GUILD_ID)).map((t) => t.threadId)).toEqual([threadId]);
expect((await repositories.threads.listArchivedThreads(FORUM_ID, false, {limit: 10})).threads).toEqual([]);
});
it('serialises racing transitions through the state version', async () => {
const threadId = freshThreadId();
await repositories.threads.create(createParams(threadId));
const [archive, lock] = await Promise.all([
repositories.threads.updateState(threadId, () => ({archived: true})),
repositories.threads.updateState(threadId, () => ({locked: true})),
]);
expect(archive).not.toBeNull();
expect(lock).not.toBeNull();
const state = await repositories.threads.getState(threadId);
expect(state?.archived).toBe(true);
expect(state?.locked).toBe(true);
expect(state?.stateVersion).toBe(3);
expect(await repositories.threads.listActiveThreads(GUILD_ID)).toEqual([]);
expect((await repositories.threads.listArchivedThreads(PARENT_ID, false, {limit: 10})).threads).toHaveLength(1);
});
it('drops and reports stale index rows, and repairs missing ones', async () => {
const drift = vi.fn();
const threads = new ThreadRepository(repositories.channelData, repositories.messages, drift);
const threadId = freshThreadId();
await threads.create(createParams(threadId));
await threads.updateState(threadId, () => ({archived: true}));
await upsertOne(
ActiveThreadsByGuild.upsertAll({
guild_id: GUILD_ID,
thread_id: threadId,
parent_id: PARENT_ID,
type: ChannelTypes.PUBLIC_THREAD,
}),
);
expect(await threads.listActiveThreads(GUILD_ID)).toEqual([]);
expect(drift).toHaveBeenCalledWith([threadId]);
expect(
await fetchMany(ActiveThreadsByGuild.selectCql({where: ActiveThreadsByGuild.where.eq('guild_id')}), {
guild_id: GUILD_ID,
}),
).toEqual([]);
const state = await threads.getState(threadId);
await executeConditional(
ArchivedThreadsByParent.conditionalDeleteByPk(
{parent_id: PARENT_ID, is_private: false, archive_timestamp: state!.archiveTimestamp!, thread_id: threadId},
{guild_id: GUILD_ID},
),
);
expect((await threads.listArchivedThreads(PARENT_ID, false, {limit: 10})).threads).toEqual([]);
await threads.repairThreadIndexes([threadId]);
expect((await threads.listArchivedThreads(PARENT_ID, false, {limit: 10})).threads).toHaveLength(1);
});
it('claims one pinned post per forum and recovers a stale claim', async () => {
const first = freshThreadId();
const second = freshThreadId();
await repositories.threads.create(createParams(first, {flags: ChannelFlags.PINNED}));
await repositories.threads.create(createParams(second));
expect(await repositories.threads.claimForumPin(FORUM_ID, first)).toBe(true);
expect(await repositories.threads.claimForumPin(FORUM_ID, first)).toBe(true);
expect(await repositories.threads.claimForumPin(FORUM_ID, second)).toBe(false);
await upsertOne(ThreadState.patchByPk({thread_id: first}, {flags: {kind: 'set', value: 0}}));
expect(await repositories.threads.claimForumPin(FORUM_ID, second)).toBe(true);
expect(await repositories.threads.getForumPin(FORUM_ID)).toBe(second);
});
it('pages archived threads newest first with before cursors and limits 2 and 100', async () => {
const ids: Array<ChannelID> = [];
for (let i = 0; i < 5; i++) {
const threadId = freshThreadId();
ids.push(threadId);
await repositories.threads.create(createParams(threadId));
await repositories.threads.updateState(threadId, () => ({
archived: true,
archive_timestamp: new Date(1_700_000_000_000 + i * 1000),
}));
}
const newestFirst = [...ids].reverse();
const first = await repositories.threads.listArchivedThreads(PARENT_ID, false, {limit: 2});
expect(first.threads.map((t) => t.threadId)).toEqual(newestFirst.slice(0, 2));
expect(first.hasMore).toBe(true);
const second = await repositories.threads.listArchivedThreads(PARENT_ID, false, {
limit: 2,
before: first.threads[1]!.archiveTimestamp!,
});
expect(second.threads.map((t) => t.threadId)).toEqual(newestFirst.slice(2, 4));
expect(second.hasMore).toBe(true);
const all = await repositories.threads.listArchivedThreads(PARENT_ID, false, {limit: 100});
expect(all.threads.map((t) => t.threadId)).toEqual(newestFirst);
expect(all.hasMore).toBe(false);
expect((await repositories.threads.listArchivedThreads(PARENT_ID, true, {limit: 100})).threads).toEqual([]);
});
it('pages past stale rows when several archived threads share one timestamp', async () => {
const tied = new Date(1_700_000_000_000);
const ids: Array<ChannelID> = [];
for (let i = 0; i < 4; i++) {
const threadId = freshThreadId();
ids.push(threadId);
await repositories.threads.create(createParams(threadId));
await repositories.threads.updateState(threadId, () => ({archived: true, archive_timestamp: tied}));
}
const older = freshThreadId();
await repositories.threads.create(createParams(older));
await repositories.threads.updateState(older, () => ({
archived: true,
archive_timestamp: new Date(tied.getTime() - 1000),
}));
for (const offset of [1n, 2n]) {
await upsertOne(
ArchivedThreadsByParent.upsertAll({
parent_id: PARENT_ID,
is_private: false,
archive_timestamp: tied,
thread_id: createChannelID(ids[3]! + 1000n * offset),
guild_id: GUILD_ID,
}),
);
}
const newestFirst = [...ids].reverse();
const first = await repositories.threads.listArchivedThreads(PARENT_ID, false, {limit: 2});
expect(first.threads.map((t) => t.threadId)).toEqual(newestFirst.slice(0, 2));
expect(first.hasMore).toBe(true);
const all = await repositories.threads.listArchivedThreads(PARENT_ID, false, {limit: 4});
expect(all.threads.map((t) => t.threadId)).toEqual(newestFirst);
expect(all.hasMore).toBe(true);
});
it('lists only joined private archived threads, newest id first', async () => {
const member = createUserID(1_900_000_000_000_002_000n);
const privateIds: Array<ChannelID> = [];
for (let i = 0; i < 3; i++) {
const threadId = freshThreadId();
privateIds.push(threadId);
const params = createParams(threadId, {members: [{userId: member, flags: 1}]});
params.channel.type = ChannelTypes.PRIVATE_THREAD;
await repositories.threads.create(params);
await repositories.threads.updateState(threadId, () => ({archived: true}));
}
const publicId = freshThreadId();
await repositories.threads.create(createParams(publicId, {members: [{userId: member, flags: 1}]}));
await repositories.threads.updateState(publicId, () => ({archived: true}));
executor.statements.length = 0;
const page = await repositories.threads.listJoinedPrivateArchivedThreads(member, GUILD_ID, PARENT_ID, {limit: 2});
expect(page.threads.map((t) => t.threadId)).toEqual([privateIds[2], privateIds[1]]);
expect(page.hasMore).toBe(true);
const next = await repositories.threads.listJoinedPrivateArchivedThreads(member, GUILD_ID, PARENT_ID, {
limit: 2,
before: privateIds[1],
});
expect(next.threads.map((t) => t.threadId)).toEqual([privateIds[0]]);
expect(next.hasMore).toBe(false);
expect(page.threads.every((t) => t.isPrivate)).toBe(true);
});
it('adds a 250 member batch with one state compare-and-set', async () => {
const threadId = freshThreadId();
await repositories.threads.create(createParams(threadId));
executor.statements.length = 0;
const members = Array.from({length: 250}, (_, index) => ({
userId: createUserID(1_900_000_000_010_000_000n + BigInt(index)),
flags: 0,
}));
const result = await repositories.threads.addMembers(threadId, members);
expect(result?.added).toHaveLength(250);
expect(result?.state.memberCount).toBe(251);
expect(result?.state.memberIdsPreview).toHaveLength(8);
expect(result?.state.memberIdsPreview[0]).toBe(members[249]!.userId);
expect(executor.count('cas:thread_state')).toBe(1);
expect(executor.count('cas:thread_members')).toBe(1);
const again = await repositories.threads.addMembers(threadId, members.slice(0, 3));
expect(again?.added).toEqual([]);
expect((await repositories.threads.listMembers(threadId, {limit: 1000})).length).toBe(251);
});
it('writes creator memberships and the seeded stamp only under compare-and-set', async () => {
const threadId = freshThreadId();
await repositories.threads.create(
createParams(threadId, {
members: [
{userId: OWNER_ID, flags: 1},
{userId: OWNER_ID, flags: 1},
],
}),
);
expect(executor.count('cas:thread_members')).toBe(1);
expect(executor.statements.some((entry) => entry.endsWith(':thread_members') && !entry.startsWith('cas:'))).toBe(
false,
);
expect((await repositories.threads.listMembers(threadId, {limit: 10})).map((m) => m.userId)).toEqual([OWNER_ID]);
const guildId = createGuildID(1_900_000_000_000_000_777n);
await repositories.threads.ensureGuildMarker(guildId);
expect((await repositories.threads.getGuildMarker(guildId))?.perms_seeded_at ?? null).toBeNull();
executor.statements.length = 0;
const seededAt = new Date(1_800_000_000_000);
await repositories.threads.markGuildPermsSeeded(guildId, seededAt);
await repositories.threads.markGuildPermsSeeded(guildId, new Date());
expect((await repositories.threads.getGuildMarker(guildId))?.perms_seeded_at).toEqual(seededAt);
expect(
executor.statements.filter((entry) => entry.endsWith(':guild_thread_state') && !entry.startsWith('select:')),
).toEqual(['cas:guild_thread_state']);
const fresh = createGuildID(1_900_000_000_000_000_778n);
await repositories.threads.markGuildPermsSeeded(fresh, seededAt);
expect((await repositories.threads.getGuildMarker(fresh))?.perms_seeded_at).toEqual(seededAt);
});
it('marks search backfill on a full marker row and clears it for a reindex', async () => {
const guildId = createGuildID(1_900_000_000_000_000_779n);
const at = new Date(1_800_000_000_000);
await repositories.threads.markGuildSearchBackfilled(guildId, at);
const marker = await repositories.threads.getGuildMarker(guildId);
expect(marker?.first_active_at).toBeInstanceOf(Date);
expect(marker?.search_backfilled_at).toEqual(at);
await repositories.threads.clearGuildSearchBackfilled(guildId);
const cleared = await repositories.threads.getGuildMarker(guildId);
expect(cleared?.search_backfilled_at ?? null).toBeNull();
expect(cleared?.first_active_at).toEqual(marker?.first_active_at);
});
it('refuses joins past the member cap', async () => {
const threadId = freshThreadId();
await repositories.threads.create(createParams(threadId));
await upsertOne(
ThreadState.patchByPk({thread_id: threadId}, {member_count: {kind: 'set', value: MAX_THREAD_MEMBERS - 1}}),
);
await expect(
repositories.threads.addMembers(threadId, [
{userId: createUserID(1_900_000_000_020_000_001n), flags: 0},
{userId: createUserID(1_900_000_000_020_000_002n), flags: 0},
]),
).rejects.toBeInstanceOf(MaxThreadMembersError);
const single = await repositories.threads.addMembers(threadId, [
{userId: createUserID(1_900_000_000_020_000_003n), flags: 0},
]);
expect(single?.state.memberCount).toBe(MAX_THREAD_MEMBERS);
});
it('lets only one of two concurrent joins take the last member slot', async () => {
const threadId = freshThreadId();
await repositories.threads.create(createParams(threadId));
await upsertOne(
ThreadState.patchByPk({thread_id: threadId}, {member_count: {kind: 'set', value: MAX_THREAD_MEMBERS - 1}}),
);
const joiners = [createUserID(1_900_000_000_025_000_001n), createUserID(1_900_000_000_025_000_002n)];
const results = await Promise.allSettled(
joiners.map((userId) => repositories.threads.addMembers(threadId, [{userId, flags: 0}])),
);
expect(results.filter((result) => result.status === 'fulfilled')).toHaveLength(1);
const rejected = results.find((result) => result.status === 'rejected') as PromiseRejectedResult;
expect(rejected.reason).toBeInstanceOf(MaxThreadMembersError);
expect((await repositories.threads.getState(threadId))?.memberCount).toBe(MAX_THREAD_MEMBERS);
expect(await repositories.threads.listMembers(threadId, {limit: 100})).toHaveLength(2);
});
it('keeps the member count consistent under concurrent adds and removes', async () => {
const threadId = freshThreadId();
const leaving = createUserID(1_900_000_000_030_000_001n);
await repositories.threads.create(
createParams(threadId, {
members: [
{userId: OWNER_ID, flags: 1},
{userId: leaving, flags: 0},
],
}),
);
const joiners = [createUserID(1_900_000_000_030_000_002n), createUserID(1_900_000_000_030_000_003n)];
await Promise.all([
repositories.threads.addMembers(
threadId,
joiners.map((userId) => ({userId, flags: 0})),
),
repositories.threads.removeMembers(threadId, [leaving]),
repositories.threads.removeMembers(threadId, [leaving]),
]);
const state = await repositories.threads.getState(threadId);
const rows = await repositories.threads.listMembers(threadId, {limit: 100});
expect(rows.map((row) => row.userId).sort()).toEqual([OWNER_ID, ...joiners].sort());
expect(state?.memberCount).toBe(3);
expect(state?.memberIdsPreview).not.toContain(leaving);
expect(await repositories.threads.listJoinedThreadIds(leaving, GUILD_ID)).toEqual([]);
});
it('rolls back member rows when the member count compare-and-set gives up', async () => {
const threadId = freshThreadId();
const staying = createUserID(1_900_000_000_040_000_001n);
const joining = createUserID(1_900_000_000_040_000_002n);
await repositories.threads.create(
createParams(threadId, {
members: [
{userId: OWNER_ID, flags: 1},
{userId: staying, flags: 0},
],
}),
);
const executeQuery = executor.executeQuery.bind(executor);
const spy = vi.spyOn(executor, 'executeQuery').mockImplementation(async (query) => {
const meta = query.kvMeta;
if (meta?.table.name === 'thread_state' && meta.conditions) {
return [{'[applied]': false}] as never;
}
return executeQuery(query);
});
await expect(repositories.threads.addMembers(threadId, [{userId: joining, flags: 0}])).rejects.toThrow();
await expect(repositories.threads.removeMembers(threadId, [staying])).rejects.toThrow();
spy.mockRestore();
const rows = await repositories.threads.listMembers(threadId, {limit: 100});
expect(rows.map((row) => row.userId).sort()).toEqual([OWNER_ID, staying].sort());
expect((await repositories.threads.getState(threadId))?.memberCount).toBe(2);
const joined = await repositories.threads.addMembers(threadId, [{userId: joining, flags: 0}]);
expect(joined?.added.map((member) => member.userId)).toEqual([joining]);
expect(joined?.state.memberCount).toBe(3);
expect(await repositories.threads.listJoinedThreadIds(joining, GUILD_ID)).toEqual([threadId]);
const left = await repositories.threads.removeMembers(threadId, [staying]);
expect(left.removed.map((member) => member.userId)).toEqual([staying]);
expect(left.state?.memberCount).toBe(2);
});
it('counts inserted thread messages only, excludes the starter id and floors at zero', async () => {
const threadId = freshThreadId();
await repositories.threads.create(createParams(threadId));
await repositories.channelData.updateLastMessageId(threadId, channelIdToMessageId(threadId), {isInsert: true});
expect((await repositories.threads.getStats(threadId)).messageCount).toBe(0);
await repositories.channelData.updateLastMessageId(threadId, createMessageID(threadId + 10n));
expect((await repositories.threads.getStats(threadId)).messageCount).toBe(0);
await repositories.channelData.updateLastMessageId(threadId, createMessageID(threadId + 20n), {isInsert: true});
await repositories.channelData.updateLastMessageId(threadId, createMessageID(threadId + 30n), {isInsert: true});
const stats = await repositories.threads.getStats(threadId);
expect(stats.messageCount).toBe(2);
expect(stats.totalMessageSent).toBe(2);
await repositories.threads.adjustMessageCount(threadId, -5);
const floored = await repositories.threads.getStats(threadId);
expect(floored.messageCount).toBe(0);
expect(floored.totalMessageSent).toBe(2);
executor.statements.length = 0;
await repositories.channelData.updateLastMessageId(PARENT_ID, createMessageID(threadId + 40n), {isInsert: true});
expect(executor.count('select:thread_stats')).toBe(0);
});
it('keeps thread message counters exact under concurrent sends and deletes', async () => {
const threadId = freshThreadId();
await repositories.threads.create(createParams(threadId));
await Promise.all(
Array.from({length: 8}, (_, index) =>
repositories.channelData.updateLastMessageId(threadId, createMessageID(threadId + BigInt(index + 1)), {
isInsert: true,
}),
),
);
const sent = await repositories.threads.getStats(threadId);
expect(sent.messageCount).toBe(8);
expect(sent.totalMessageSent).toBe(8);
await Promise.all([
repositories.threads.adjustMessageCount(threadId, -1),
repositories.threads.adjustMessageCount(threadId, -1),
repositories.channelData.updateLastMessageId(threadId, createMessageID(threadId + 100n), {isInsert: true}),
]);
const mixed = await repositories.threads.getStats(threadId);
expect(mixed.messageCount).toBe(7);
expect(mixed.totalMessageSent).toBe(9);
});
it('enumerates guild threads with and without an activity cutoff', async () => {
await taint();
const oldThread = freshThreadId(-86_400_000);
const newThread = freshThreadId();
await repositories.threads.create(createParams(oldThread));
await repositories.threads.create(createParams(newThread));
await repositories.channelData.updateLastMessageId(
newThread,
createMessageID(createSnowflakeFromTimestamp(Date.now())),
);
expect((await repositories.threads.listGuildThreadIds(GUILD_ID)).sort()).toEqual([oldThread, newThread].sort());
expect(
await repositories.threads.listGuildThreadIds(GUILD_ID, {activeSince: new Date(Date.now() - 3_600_000)}),
).toEqual([newThread]);
});
it('does no thread IO for never-enabled guilds', async () => {
const threadId = freshThreadId();
await upsertOne(
ThreadsByParent.upsertAll({parent_id: PARENT_ID, thread_id: threadId, guild_id: GUILD_ID, type: 11}),
);
executor.statements.length = 0;
expect(await repositories.threads.listGuildThreadIds(GUILD_ID)).toEqual([]);
expect(executor.statements).toEqual([]);
await repositories.channelData.listGuildChannels(GUILD_ID, 'maintenance');
await repositories.channelData.listGuildChannels(GUILD_ID, 'enrolled');
expect(executor.statements.every((entry) => /:(channels|channels_by_guild_id)$/.test(entry))).toBe(true);
expect(executor.count('select:guild_thread_state')).toBe(0);
});
it('indexes forums separately and merges them only for active or tainted guilds', async () => {
await repositories.channelData.upsert(channelRow(FORUM_ID, ChannelTypes.GUILD_FORUM));
expect(
await fetchMany(ThreadOnlyChannelsByGuild.selectCql({where: ThreadOnlyChannelsByGuild.where.eq('guild_id')}), {
guild_id: GUILD_ID,
}),
).toHaveLength(1);
expect((await repositories.channelData.listGuildChannels(GUILD_ID, 'enrolled')).map((c) => c.id)).toEqual([
PARENT_ID,
]);
expect((await repositories.channelData.listGuildChannels(GUILD_ID, 'maintenance')).map((c) => c.id)).toEqual([
PARENT_ID,
]);
await taint();
expect(
new Set((await repositories.channelData.listGuildChannels(GUILD_ID, 'enrolled')).map((c) => c.id)),
).toEqual(new Set([PARENT_ID, FORUM_ID]));
setThreadsConfig({enabled: false, ever_enabled: true});
expect((await repositories.channelData.listGuildChannels(GUILD_ID, 'enrolled')).map((c) => c.id)).toEqual([
PARENT_ID,
]);
expect(
new Set((await repositories.channelData.listGuildChannels(GUILD_ID, 'maintenance')).map((c) => c.id)),
).toEqual(new Set([PARENT_ID, FORUM_ID]));
await repositories.channelData.delete(FORUM_ID, GUILD_ID, ChannelTypes.GUILD_FORUM);
expect(
await fetchMany(ThreadOnlyChannelsByGuild.selectCql({where: ThreadOnlyChannelsByGuild.where.eq('guild_id')}), {
guild_id: GUILD_ID,
}),
).toEqual([]);
});
it('purges every thread row with threads_by_parent last', async () => {
const threadId = freshThreadId();
await repositories.threads.create(createParams(threadId, {members: [{userId: OWNER_ID, flags: 1}]}));
await repositories.threads.updateState(threadId, () => ({archived: true}));
executor.statements.length = 0;
await repositories.threads.purgeThread(threadId);
expect(executor.statements.at(-1)).toBe('delete:threads_by_parent');
expect(await repositories.threads.getState(threadId)).toBeNull();
expect(await repositories.channelData.findUnique(threadId)).toBeNull();
expect(await repositories.threads.listMembers(threadId, {limit: 10})).toEqual([]);
expect(await repositories.threads.listJoinedThreadIds(OWNER_ID, GUILD_ID)).toEqual([]);
expect((await repositories.threads.listArchivedThreads(PARENT_ID, false, {limit: 10})).threads).toEqual([]);
expect(await repositories.threads.listThreadIdsByParent(PARENT_ID, {limit: 10})).toEqual([]);
});
it('stores parent config through plain patches', async () => {
await repositories.threads.patchParentConfig(GUILD_ID, FORUM_ID, {
flags: ChannelFlags.REQUIRE_TAG,
available_tags: [{id: 5n, name: 'bug', moderated: false, emoji_id: null, emoji_name: '🐛'}],
default_sort_order: 1,
});
const config = await repositories.threads.getParentConfig(GUILD_ID, FORUM_ID);
expect(config?.flags).toBe(ChannelFlags.REQUIRE_TAG);
expect(config?.availableTags.map((tag) => tag.toUdt())).toEqual([
{id: 5n, name: 'bug', moderated: false, emoji_id: null, emoji_name: '🐛'},
]);
await repositories.threads.patchParentConfig(GUILD_ID, FORUM_ID, {available_tags: []});
expect((await repositories.threads.getParentConfig(GUILD_ID, FORUM_ID))?.availableTags).toEqual([]);
expect(await repositories.threads.listParentConfigs(GUILD_ID)).toHaveLength(1);
await repositories.threads.deleteParentConfig(GUILD_ID, FORUM_ID);
expect(
await fetchOne(
ThreadParentConfig.selectCql({
where: [ThreadParentConfig.where.eq('guild_id'), ThreadParentConfig.where.eq('channel_id')],
}),
{guild_id: GUILD_ID, channel_id: FORUM_ID},
),
).toBeNull();
});
});
}
const THREAD_TABLES = [
ThreadState,
ThreadStats,
ThreadsByParent,
ActiveThreadsByGuild,
ArchivedThreadsByParent,
ThreadMembers,
ThreadMembersByUser,
ThreadParentConfig,
ForumPinnedThread,
ThreadOnlyChannelsByGuild,
GuildThreadState,
];
describe('thread storage leaves control storage untouched', () => {
it('keeps the channel and message column lists byte-identical', () => {
expect([...CHANNEL_COLUMNS]).toEqual([
'channel_id',
'guild_id',
'type',
'name',
'topic',
'icon_hash',
'url',
'parent_id',
'position',
'owner_id',
'recipient_ids',
'nsfw',
'content_warning_level',
'content_warning_text',
'rate_limit_per_user',
'bitrate',
'user_limit',
'voice_connection_limit',
'rtc_region',
'last_message_id',
'last_pin_timestamp',
'permission_overwrites',
'nicks',
'soft_deleted',
'indexed_at',
'version',
]);
expect([...MESSAGE_COLUMNS]).toEqual([
'channel_id',
'bucket',
'message_id',
'author_id',
'type',
'webhook_id',
'webhook_name',
'webhook_avatar_hash',
'content',
'edited_timestamp',
'pinned_timestamp',
'flags',
'mention_everyone',
'mention_users',
'mention_roles',
'mention_channels',
'attachments',
'embeds',
'sticker_items',
'message_reference',
'message_snapshots',
'call',
'has_reaction',
'version',
]);
});
it('declares every thread table in the target schema with the same columns and no TTL', () => {
const schema = JSON.parse(
readFileSync(
fileURLToPath(new URL('../../../../../tools/dev/cassandra_target_schema.json', import.meta.url)),
'utf8',
),
) as {tables: Array<{name: string; columns: Array<{name: string}>; options: string}>};
for (const table of THREAD_TABLES) {
const declared = schema.tables.find((entry) => entry.name === table.name);
expect(declared, table.name).toBeDefined();
expect(new Set(declared!.columns.map((column) => column.name))).toEqual(new Set(table.columns));
expect(declared!.options).not.toContain('default_time_to_live');
expect(table.defaultTtlSeconds).toBeUndefined();
}
});
});
describeThreadRepository('in-memory cassandra', async () => new InMemoryCassandraQueryExecutor());
const dockerAvailable = spawnSync('docker', ['version'], {stdio: 'ignore'}).status === 0;
const KV_TABLE = 'kv_thread_repository';
const CONTAINER = `fluxer-threads-${process.pid.toString(36)}-${Date.now().toString(36)}`;
async function freePort(): Promise<number> {
return new Promise((resolve, reject) => {
const server = createServer();
server.on('error', reject);
server.listen(0, '127.0.0.1', () => {
const address = server.address();
if (typeof address === 'string' || address === null) {
reject(new Error('no port'));
return;
}
const port = address.port;
server.close(() => resolve(port));
});
});
}
describe.skipIf(!dockerAvailable)('ThreadRepository against postgres kv', () => {
beforeAll(async () => {
const port = await freePort();
startDockerContainer([
'run',
'-d',
'--name',
CONTAINER,
'-e',
'POSTGRES_USER=fluxer',
'-e',
'POSTGRES_PASSWORD=fluxer',
'-e',
'POSTGRES_DB=fluxer',
'-p',
`127.0.0.1:${port}:5432`,
'postgres:16-alpine',
'-c',
'fsync=off',
]);
let ready = false;
for (let attempt = 0; attempt < 180 && !ready; attempt += 1) {
await new Promise((resolve) => setTimeout(resolve, 500));
const probe = spawnSync('docker', ['exec', CONTAINER, 'pg_isready', '-U', 'fluxer', '-d', 'fluxer'], {
stdio: 'ignore',
});
if (probe.status !== 0) continue;
try {
await initPostgres({
url: `postgres://fluxer:[email protected]:${port}/fluxer`,
maxConnections: 8,
kvTable: KV_TABLE,
});
await getDefaultPostgresClient().query('SELECT 1');
ready = true;
} catch {
await shutdownPostgres().catch(() => {});
}
}
if (!ready) throw new Error('postgres never came up');
await ensurePostgresKvSchema(getDefaultPostgresClient());
}, 900_000);
afterAll(async () => {
setCassandraQueryExecutorForTesting(new InMemoryCassandraQueryExecutor());
await shutdownPostgres().catch(() => {});
spawnSync('docker', ['rm', '-f', CONTAINER], {stdio: 'ignore'});
});
describeThreadRepository('postgres kv', async () => {
const client = getDefaultPostgresClient();
await client.query(`DELETE FROM ${KV_TABLE}`);
return new PostgresKvQueryExecutor(client);
});
});
File diff suppressed because it is too large Load Diff
@@ -22,6 +22,7 @@ import {
} from '@app/api/channel/services/message/MessageHelpers';
import {applyUploadRelayDecision, resolveUploadRelayDecision} from '@app/api/channel/services/UploadRelay';
import {SYSTEM_USER_ID} from '@app/api/constants/Core';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import type {IPurgeQueue} from '@app/api/infrastructure/CachePurgeQueue';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import type {IStorageService} from '@app/api/infrastructure/IStorageService';
@@ -44,6 +45,7 @@ import {
ATTACHMENT_UPLOAD_MAX_CHUNKS,
resolveAttachmentUploadPartSize,
} from '@fluxer/constants/src/LimitConstants';
import {THREAD_FEATURE_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {ValidationErrorCodes} from '@fluxer/constants/src/ValidationErrorCodes';
import {CannotSendMessageToNonTextChannelError} from '@fluxer/errors/src/domains/channel/CannotSendMessageToNonTextChannelError';
import {UnknownChannelError} from '@fluxer/errors/src/domains/channel/UnknownChannelError';
@@ -64,6 +66,7 @@ import type {
interface DeleteAttachmentParams {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
attachmentId: AttachmentID;
@@ -74,6 +77,7 @@ type UploadActor = 'member' | 'webhook';
interface UploadFormDataAttachmentsParams {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
clientIp: string;
files: Array<{
@@ -89,6 +93,7 @@ interface UploadFormDataAttachmentsParams {
interface RequestPresignedAttachmentUploadUrlsParams {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
clientIp: string;
attachments: Array<PresignedAttachmentUploadRequestItem>;
@@ -96,6 +101,7 @@ interface RequestPresignedAttachmentUploadUrlsParams {
interface CompleteMultipartAttachmentUploadsParams {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
clientIp: string;
uploads: Array<CompleteMultipartAttachmentUploadItem>;
@@ -118,13 +124,14 @@ export class AttachmentUploadService {
async uploadFormDataAttachments({
userId,
viewer,
channelId,
clientIp,
files,
attachmentMetadata,
actor = 'member',
}: UploadFormDataAttachmentsParams): Promise<Array<UploadedAttachment>> {
const {maxFileSize} = await this.getUploadPermissionAndLimit({userId, channelId, actor});
const {maxFileSize} = await this.getUploadPermissionAndLimit({userId, viewer, channelId, actor});
assertAttachmentFileSizesWithinLimit(
files.map(({file}) => file.size),
maxFileSize,
@@ -172,6 +179,7 @@ export class AttachmentUploadService {
async requestPresignedAttachmentUploadUrls({
userId,
viewer,
channelId,
clientIp,
attachments,
@@ -179,7 +187,7 @@ export class AttachmentUploadService {
if (!Config.presignedAttachmentUploadsEnabled) {
throw new FeatureTemporarilyDisabledError();
}
const {maxFileSize} = await this.getUploadPermissionAndLimit({userId, channelId, actor: 'member'});
const {maxFileSize} = await this.getUploadPermissionAndLimit({userId, viewer, channelId, actor: 'member'});
assertAttachmentFileSizesWithinLimit(
attachments.map(({file_size}) => file_size),
maxFileSize,
@@ -279,6 +287,7 @@ export class AttachmentUploadService {
async completeMultipartAttachmentUploads({
userId,
viewer,
channelId,
clientIp,
uploads,
@@ -286,7 +295,7 @@ export class AttachmentUploadService {
if (!Config.presignedAttachmentUploadsEnabled) {
throw new FeatureTemporarilyDisabledError();
}
const {maxFileSize} = await this.getUploadPermissionAndLimit({userId, channelId, actor: 'member'});
const {maxFileSize} = await this.getUploadPermissionAndLimit({userId, viewer, channelId, actor: 'member'});
const bucket = Config.s3.buckets.uploads;
return Promise.all(
uploads.map(async ({upload_filename, upload_id}, index) => {
@@ -349,6 +358,7 @@ export class AttachmentUploadService {
async deleteAttachment({
userId,
viewer,
channelId,
messageId,
attachmentId,
@@ -357,6 +367,7 @@ export class AttachmentUploadService {
const {channel, guild} = await this.messageInteractionService.authService.getChannelAuthenticated({
userId,
channelId,
viewer,
});
if (isOperationDisabled(guild, GuildOperations.SEND_MESSAGE)) {
throw new FeatureTemporarilyDisabledError();
@@ -380,6 +391,7 @@ export class AttachmentUploadService {
if (willBeEmpty) {
await this.messageService.deletion.deleteMessage({
userId,
viewer,
channelId,
messageId,
requestCache,
@@ -411,6 +423,7 @@ export class AttachmentUploadService {
if (!updatedMessage) {
await this.messageService.deletion.deleteMessage({
userId,
viewer,
channelId,
messageId,
requestCache,
@@ -449,10 +462,12 @@ export class AttachmentUploadService {
private async getUploadPermissionAndLimit({
userId,
viewer,
channelId,
actor,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
actor: UploadActor;
}): Promise<{
@@ -461,8 +476,8 @@ export class AttachmentUploadService {
const {channel, guild} =
actor === 'webhook'
? await this.getWebhookUploadChannel(channelId)
: await this.getMemberUploadChannel({userId, channelId});
if (!TEXT_BASED_CHANNEL_TYPES.has(channel.type)) {
: await this.getMemberUploadChannel({userId, viewer, channelId});
if (!TEXT_BASED_CHANNEL_TYPES.has(channel.type) && !THREAD_FEATURE_CHANNEL_TYPES.has(channel.type)) {
throw new CannotSendMessageToNonTextChannelError();
}
const user = await this.userRepository.findUnique(userId);
@@ -478,7 +493,15 @@ export class AttachmentUploadService {
return {maxFileSize};
}
private async getMemberUploadChannel({userId, channelId}: {userId: UserID; channelId: ChannelID}): Promise<{
private async getMemberUploadChannel({
userId,
viewer,
channelId,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
}): Promise<{
channel: Channel;
guild: GuildResponse | null;
}> {
@@ -486,6 +509,7 @@ export class AttachmentUploadService {
await this.messageInteractionService.authService.getChannelAuthenticated({
userId,
channelId,
viewer,
});
if (guild) {
await checkPermission(Permissions.SEND_MESSAGES | Permissions.ATTACH_FILES);
@@ -1,13 +1,27 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {Channel} from '@app/api/models/Channel';
import type {ThreadMember} from '@app/api/models/ThreadMember';
import type {ThreadState} from '@app/api/models/ThreadState';
import type {ThreadActorContext} from '@fluxer/constants/src/ThreadPermissionUtils';
import type {GuildMemberResponse} from '@fluxer/schema/src/domains/guild/GuildMemberSchemas';
import type {GuildResponse} from '@fluxer/schema/src/domains/guild/GuildResponseSchemas';
export interface AuthenticatedThread {
state: ThreadState;
parent: Channel;
member: ThreadMember | null;
parentPermissions: bigint;
actor: ThreadActorContext;
isModerator: boolean;
enforceMfa: (permission: bigint) => void;
}
export interface AuthenticatedChannel {
channel: Channel;
guild: GuildResponse | null;
member: GuildMemberResponse | null;
hasPermission: (permission: bigint) => Promise<boolean>;
checkPermission: (permission: bigint) => Promise<void>;
thread?: AuthenticatedThread;
}
@@ -2,12 +2,13 @@
import type {ChannelID, GuildID, UserID} from '@app/api/BrandedTypes';
import type {IChannelRepositoryAggregate} from '@app/api/channel/repositories/IChannelRepositoryAggregate';
import type {AuthenticatedChannel} from '@app/api/channel/services/AuthenticatedChannel';
import type {AuthenticatedChannel, AuthenticatedThread} from '@app/api/channel/services/AuthenticatedChannel';
import {DMPermissionValidator} from '@app/api/channel/services/DMPermissionValidator';
import {
ensurePersonalNotesChannelExists,
isPersonalNotesChannelId,
} from '@app/api/channel/services/PersonalNotesChannelRepair';
import {assertThreadAllowed} from '@app/api/channel/services/thread/ThreadDenials';
import {
type ContentWarningChannelLike,
channelResponseToContentWarningView,
@@ -16,6 +17,15 @@ import {
guildResponseToContentWarningView,
} from '@app/api/channel/utils/EffectiveContentWarning';
import {SYSTEM_USER_ID} from '@app/api/constants/Core';
import {
THREAD_CHANNEL_TYPES,
THREAD_FEATURE_CHANNEL_TYPES,
THREAD_ONLY_CHANNEL_TYPES,
THREAD_PARENT_CHANNEL_TYPES,
type ThreadViewer,
viewerActive,
} from '@app/api/experiment/ChannelThreadsGate';
import {isGuildMemberTimedOut} from '@app/api/guild/GuildModel';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
import {createGuildMfaEnforcer} from '@app/api/guild/services/GuildMfaEnforcement';
import type {GuildChannelAuthContext, IGatewayService} from '@app/api/infrastructure/IGatewayService';
@@ -25,6 +35,12 @@ import type {User} from '@app/api/models/User';
import type {IUserRepository} from '@app/api/user/IUserRepository';
import {canUserAccessNsfwContent} from '@app/api/utils/AgeUtils';
import {ChannelTypes, Permissions} from '@fluxer/constants/src/ChannelConstants';
import {
canViewThread,
isThreadModerator,
threadViewPermissions,
withImplicitThreadBits,
} from '@fluxer/constants/src/ThreadPermissionUtils';
import {CannotSendMessagesToUserError} from '@fluxer/errors/src/domains/channel/CannotSendMessagesToUserError';
import {UnknownChannelError} from '@fluxer/errors/src/domains/channel/UnknownChannelError';
import {AccessDeniedError} from '@fluxer/errors/src/domains/core/AccessDeniedError';
@@ -33,6 +49,7 @@ import {UnknownGuildError} from '@fluxer/errors/src/domains/guild/UnknownGuildEr
import {NsfwContentRequiresAgeVerificationError} from '@fluxer/errors/src/domains/moderation/NsfwContentRequiresAgeVerificationError';
import {UnknownUserError} from '@fluxer/errors/src/domains/user/UnknownUserError';
import type {GuildMemberResponse} from '@fluxer/schema/src/domains/guild/GuildMemberSchemas';
import type {GuildResponse} from '@fluxer/schema/src/domains/guild/GuildResponseSchemas';
export interface ChannelAuthOptions {
errorOnMissingGuild: 'unknown_channel' | 'missing_permissions';
@@ -51,6 +68,13 @@ interface DMSendPermissionsByChannelIdParams {
type DMSendPermissionsParams = DMSendPermissionsByChannelParams | DMSendPermissionsByChannelIdParams;
export interface ThreadPermissionContext {
guild: GuildResponse;
member: GuildMemberResponse;
parentCategory: GuildChannelAuthContext['parentChannel'];
thread: AuthenticatedThread;
}
export abstract class BaseChannelAuthService {
protected abstract readonly options: ChannelAuthOptions;
protected dmPermissionValidator: DMPermissionValidator;
@@ -70,10 +94,12 @@ export abstract class BaseChannelAuthService {
async getChannelAuthenticated({
userId,
channelId,
viewer,
skipNsfwValidation,
}: {
userId: UserID;
channelId: ChannelID;
viewer: ThreadViewer;
skipNsfwValidation?: boolean;
}): Promise<AuthenticatedChannel> {
if (this.isPersonalNotesChannel({userId, channelId})) {
@@ -85,6 +111,12 @@ export abstract class BaseChannelAuthService {
}
const channel = await this.channelRepository.channelData.findUnique(channelId);
if (!channel) throw new UnknownChannelError();
if (
THREAD_FEATURE_CHANNEL_TYPES.has(channel.type) &&
(channel.guildId === null || !viewerActive(viewer, channel.guildId))
) {
throw new UnknownChannelError();
}
if (!channel.guildId) {
const recipients = await this.userRepository.listUsers(Array.from(channel.recipientIds));
return this.getDMChannelAuth({channel, recipients, userId});
@@ -174,6 +206,9 @@ export abstract class BaseChannelAuthService {
userId: UserID;
skipNsfwValidation?: boolean;
}): Promise<AuthenticatedChannel> {
if (THREAD_CHANNEL_TYPES.has(channel.type)) {
return this.getThreadChannelAuth({channel, userId, skipNsfwValidation});
}
const guildId = channel.guildId!;
const [authContextResult, guildMemberResult] = await Promise.all([
this.fetchGuildAuthContextOrThrow({guildId, userId, channelId: this.parentLookupChannelId(channel)}),
@@ -226,7 +261,8 @@ export abstract class BaseChannelAuthService {
(channel.type === ChannelTypes.GUILD_TEXT ||
channel.type === ChannelTypes.GUILD_ANNOUNCEMENT ||
channel.type === ChannelTypes.GUILD_VOICE ||
channel.type === ChannelTypes.GUILD_LINK) &&
channel.type === ChannelTypes.GUILD_LINK ||
THREAD_ONLY_CHANNEL_TYPES.has(channel.type)) &&
requiresAgeVerification
) {
const user = await this.userRepository.findUnique(userId);
@@ -244,6 +280,111 @@ export abstract class BaseChannelAuthService {
};
}
async resolveThreadPermissionContext({
thread,
userId,
}: {
thread: Channel;
userId: UserID;
}): Promise<ThreadPermissionContext> {
const guildId = thread.guildId;
const parentId = thread.parentId;
if (guildId === null || parentId === null) throw new UnknownChannelError();
const [parent, state, threadMember, guildMemberResult] = await Promise.all([
this.channelRepository.channelData.findUnique(parentId),
this.channelRepository.threads.getState(thread.id),
this.channelRepository.threads.getMember(thread.id, userId),
this.fetchGuildMemberOrThrow({guildId, userId}),
]);
if (!parent || !state || parent.guildId !== guildId || !THREAD_PARENT_CHANNEL_TYPES.has(parent.type)) {
throw new UnknownChannelError();
}
if (!guildMemberResult.success || !guildMemberResult.memberData) {
this.throwGuildAccessError();
}
const [authContextResult, parentPermissions] = await Promise.all([
this.fetchGuildAuthContextOrThrow({guildId, userId, channelId: this.parentLookupChannelId(parent)}),
this.gatewayService.getUserPermissions({guildId, userId, channelId: parent.id}),
]);
if (!authContextResult) {
this.throwGuildAccessError();
}
const member = await this.fillMissingMemberTimeout({guildId, userId, memberData: guildMemberResult.memberData});
const guild = authContextResult.guild;
const enforceMfa = await createGuildMfaEnforcer({userRepository: this.userRepository, guildData: guild, userId});
const isOwner = guild.owner_id === userId.toString();
const timedOut = isGuildMemberTimedOut(member);
const actor = {
permissions: parentPermissions,
isOwner,
timedOut,
thread: {
type: state.type,
archived: state.archived,
locked: state.locked,
invitable: state.invitable ?? true,
},
isThreadOwner: thread.ownerId === userId,
isMember: threadMember !== null,
};
return {
guild,
member,
parentCategory: authContextResult.parentChannel,
thread: {
state,
parent,
member: threadMember,
parentPermissions,
actor,
isModerator: isThreadModerator(withImplicitThreadBits(parentPermissions), {isOwner, timedOut}),
enforceMfa,
},
};
}
protected async getThreadChannelAuth({
channel,
userId,
skipNsfwValidation,
}: {
channel: Channel;
userId: UserID;
skipNsfwValidation?: boolean;
}): Promise<AuthenticatedChannel> {
const context = await this.resolveThreadPermissionContext({thread: channel, userId});
const {guild, member, thread} = context;
assertThreadAllowed(canViewThread(thread.actor));
const permissions = threadViewPermissions(thread.parentPermissions);
const hasPermission = async (permission: bigint): Promise<boolean> => {
const allowed = (permissions & permission) === permission;
if (allowed) thread.enforceMfa(permission);
return allowed;
};
const checkPermission = async (permission: bigint): Promise<void> => {
if (!(await hasPermission(permission))) throw new MissingPermissionsError();
};
if (this.options.validateNsfw && !skipNsfwValidation && THREAD_PARENT_CHANNEL_TYPES.has(thread.parent.type)) {
const parentCategory = await this.getParentCategoryContentWarningView({
channel: thread.parent,
parentChannel: context.parentCategory,
});
const requiresAgeVerification = computeEffectiveChannelNsfw(
channelToContentWarningView(thread.parent),
parentCategory,
guildResponseToContentWarningView(guild),
);
if (requiresAgeVerification) {
const user = await this.userRepository.findUnique(userId);
if (!user) throw new UnknownUserError();
if (!canUserAccessNsfwContent(user)) {
throw new NsfwContentRequiresAgeVerificationError();
}
}
}
return {channel, guild, member, hasPermission, checkPermission, thread};
}
private parentLookupChannelId(channel: Channel): ChannelID | undefined {
if (!channel.parentId || channel.type === ChannelTypes.GUILD_CATEGORY) {
return undefined;
@@ -12,6 +12,9 @@ import {ChannelOperationsService} from '@app/api/channel/services/channel_data/C
import {ChannelUtilsService} from '@app/api/channel/services/channel_data/ChannelUtilsService';
import {GroupDmUpdateService} from '@app/api/channel/services/channel_data/GroupDmUpdateService';
import type {MessagePersistenceService} from '@app/api/channel/services/message/MessagePersistenceService';
import {ThreadModifyService} from '@app/api/channel/services/thread/ThreadModifyService';
import {pickThreadParentSettings} from '@app/api/channel/services/thread/ThreadParentSettings';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import type {GuildAuditLogService} from '@app/api/guild/GuildAuditLogService';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
import type {AvatarService} from '@app/api/infrastructure/AvatarService';
@@ -30,12 +33,16 @@ import type {IUserRepository} from '@app/api/user/IUserRepository';
import type {VoiceAvailabilityService} from '@app/api/voice/VoiceAvailabilityService';
import type {IWebhookRepository} from '@app/api/webhook/IWebhookRepository';
import {ChannelTypes} from '@fluxer/constants/src/ChannelConstants';
import type {ChannelUpdateRequest} from '@fluxer/schema/src/domains/channel/ChannelRequestSchemas';
import type {
ChannelUpdateGatedRequest,
ChannelUpdateNonThreadRequest,
ChannelUpdateThreadRequest,
} from '@fluxer/schema/src/domains/channel/ChannelRequestSchemas';
import type {ICacheService} from '@pkgs/cache/src/ICacheService';
import type {IRateLimitService} from '@pkgs/rate_limit/src/IRateLimitService';
type GuildChannelUpdateRequest = Exclude<
ChannelUpdateRequest,
ChannelUpdateNonThreadRequest,
{
type: typeof ChannelTypes.GROUP_DM;
}
@@ -47,6 +54,7 @@ export class ChannelDataService {
public readonly operations: ChannelOperationsService;
public readonly groupDmUpdate: GroupDmUpdateService;
public readonly utils: ChannelUtilsService;
public readonly threadModify: ThreadModifyService;
constructor(
channelRepository: IChannelRepositoryAggregate,
@@ -93,7 +101,18 @@ export class ChannelDataService {
limitConfigService,
rateLimitService,
cacheService,
snowflakeService,
);
this.threadModify = new ThreadModifyService({
channelRepository,
gatewayService,
guildAuditLogService,
rateLimitService,
cacheService,
snowflakeService,
messagePersistence: messagePersistenceService,
utils: this.utils,
});
this.groupDmUpdate = new GroupDmUpdateService(
channelRepository,
userRepository,
@@ -104,8 +123,26 @@ export class ChannelDataService {
);
}
async deleteThread({
userId,
viewer,
channelId,
requestCache,
auditLogReason,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
requestCache: RequestCache;
auditLogReason: string | null;
}): Promise<void> {
const authChannel = await this.auth.getChannelAuthenticated({userId, channelId, viewer, skipNsfwValidation: true});
await this.threadModify.deleteThread({authChannel, userId, requestCache, auditLogReason});
}
async editChannel({
userId,
viewer,
channelId,
data,
clientFeatures,
@@ -114,22 +151,34 @@ export class ChannelDataService {
typeConversion,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
data: Omit<ChannelUpdateRequest, 'type'>;
data: Omit<ChannelUpdateGatedRequest, 'type'>;
clientFeatures: ReadonlySet<string>;
requestCache: RequestCache;
auditLogReason: string | null;
typeConversion?: ChannelTypeConversion | null;
}): Promise<Channel> {
const {channel} = await this.auth.getChannelAuthenticated({userId, channelId, skipNsfwValidation: true});
const authChannel = await this.auth.getChannelAuthenticated({userId, channelId, viewer, skipNsfwValidation: true});
const {channel} = authChannel;
if (authChannel.thread) {
return this.threadModify.updateThread({
authChannel,
userId,
data: data as Omit<ChannelUpdateThreadRequest, 'type'>,
requestCache,
auditLogReason,
});
}
if (channel.type === ChannelTypes.GROUP_DM) {
const groupDmData = data as Omit<Extract<ChannelUpdateNonThreadRequest, {type: 3}>, 'type'>;
return await this.groupDmUpdate.updateGroupDmChannel({
userId,
channelId,
name: data.name !== undefined ? data.name : undefined,
icon: data.icon !== undefined ? data.icon : undefined,
ownerId: data.owner_id ? createUserID(data.owner_id) : undefined,
nicks: data.nicks,
name: groupDmData.name !== undefined ? groupDmData.name : undefined,
icon: groupDmData.icon !== undefined ? groupDmData.icon : undefined,
ownerId: groupDmData.owner_id ? createUserID(groupDmData.owner_id) : undefined,
nicks: groupDmData.nicks,
requestCache,
});
}
@@ -188,8 +237,10 @@ export class ChannelDataService {
}
return this.operations.editChannel({
userId,
viewer,
channelId,
data: channelUpdateData,
threadParent: pickThreadParentSettings(channel.type, guildChannelData),
clientFeatures,
requestCache,
auditLogReason,
@@ -21,11 +21,13 @@ export async function withChannelFollowLock<T>(
cacheService: ICacheService,
channelId: ChannelID,
fn: () => Promise<T>,
keepAliveTtlSeconds?: number,
): Promise<T> {
const ttlSeconds = keepAliveTtlSeconds ?? CHANNEL_FOLLOW_LOCK_TTL_SECONDS;
const lockKey = `channel-follow:${channelId}`;
let lockToken: string | null = null;
for (let attempt = 0; attempt < CHANNEL_FOLLOW_LOCK_ACQUIRE_ATTEMPTS; attempt++) {
lockToken = await cacheService.acquireLock(lockKey, CHANNEL_FOLLOW_LOCK_TTL_SECONDS);
lockToken = await cacheService.acquireLock(lockKey, ttlSeconds);
if (lockToken) break;
await new Promise((resolve) => setTimeout(resolve, CHANNEL_FOLLOW_LOCK_RETRY_DELAY_MS * (attempt + 1)));
}
@@ -36,10 +38,21 @@ export async function withChannelFollowLock<T>(
data: {retry_after: 1},
});
}
const token = lockToken;
const keepAlive =
keepAliveTtlSeconds === undefined
? null
: setInterval(
() => {
cacheService.extendLock(lockKey, token, keepAliveTtlSeconds).catch(() => {});
},
(keepAliveTtlSeconds * 1000) / 3,
);
try {
return await fn();
} finally {
await cacheService.releaseLock(lockKey, lockToken).catch(() => {});
if (keepAlive) clearInterval(keepAlive);
await cacheService.releaseLock(lockKey, token).catch(() => {});
}
}
@@ -6,6 +6,7 @@ import type {GatewayDispatchEvent} from '@app/api/constants/Gateway';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import type {Channel} from '@app/api/models/Channel';
import {ChannelTypes} from '@fluxer/constants/src/ChannelConstants';
import {THREAD_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
interface DispatchChannelEventParams {
gatewayService: IGatewayService;
@@ -14,6 +15,14 @@ interface DispatchChannelEventParams {
data: unknown;
}
export function withThreadContext(channel: Channel, data: unknown): unknown {
if (!THREAD_CHANNEL_TYPES.has(channel.type) || typeof data !== 'object' || data === null) return data;
return {
...data,
_fluxer_thread: {id: channel.id.toString(), parent_id: channel.parentId?.toString() ?? null, type: channel.type},
};
}
export async function dispatchChannelEvent({
gatewayService,
channel,
@@ -32,7 +41,7 @@ export async function dispatchChannelEvent({
});
}
if (channel.guildId) {
return gatewayService.dispatchGuild({guildId: channel.guildId, event, data});
return gatewayService.dispatchGuild({guildId: channel.guildId, event, data: withThreadContext(channel, data)});
}
await Promise.all(
Array.from(channel.recipientIds)
@@ -4,38 +4,94 @@ import type {ChannelID, UserID} from '@app/api/BrandedTypes';
import {mapChannelToResponse} from '@app/api/channel/ChannelMappers';
import type {ChannelService} from '@app/api/channel/services/ChannelService';
import type {ChannelTypeConversion} from '@app/api/channel/services/channel_data/ChannelOperationsService';
import {applyForumTagEdit, type ForumTagEdit} from '@app/api/channel/services/thread/ForumTagService';
import {mapThreadToResponse} from '@app/api/channel/services/thread/ThreadMappers';
import {withThreadParentFields} from '@app/api/channel/services/thread/ThreadParentSettings';
import {loadThreadViews} from '@app/api/channel/services/thread/ThreadViews';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import {maskChannelResponseThreadBits} from '@app/api/guild/services/ThreadPermissionBits';
import type {UserCacheService} from '@app/api/infrastructure/UserCacheService';
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
import type {Channel} from '@app/api/models/Channel';
import type {User} from '@app/api/models/User';
import {ChannelTypes} from '@fluxer/constants/src/ChannelConstants';
import type {ChannelUpdateRequest} from '@fluxer/schema/src/domains/channel/ChannelRequestSchemas';
import {APIErrorCodes} from '@fluxer/constants/src/ApiErrorCodes';
import {ChannelTypes, Permissions} from '@fluxer/constants/src/ChannelConstants';
import {InvalidChannelTypeError} from '@fluxer/errors/src/domains/channel/InvalidChannelTypeError';
import {UnknownChannelError} from '@fluxer/errors/src/domains/channel/UnknownChannelError';
import {ThrottledError} from '@fluxer/errors/src/domains/core/ThrottledError';
import type {ChannelUpdateGatedRequest} from '@fluxer/schema/src/domains/channel/ChannelRequestSchemas';
import type {ChannelResponse, ChannelSlowmodeStateResponse} from '@fluxer/schema/src/domains/channel/ChannelSchemas';
import type {ICacheService} from '@pkgs/cache/src/ICacheService';
const FORUM_TAGS_LOCK_TTL_SECONDS = 5;
const FORUM_TAGS_LOCK_ACQUIRE_ATTEMPTS = 6;
const FORUM_TAGS_LOCK_RETRY_DELAY_MS = 50;
export class ChannelRequestService {
constructor(
private readonly channelService: ChannelService,
private readonly userCacheService: UserCacheService,
private readonly cacheService: ICacheService,
) {}
async getChannelResponse(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
requestCache: RequestCache;
}): Promise<ChannelResponse> {
const channel = await this.channelService.channelData.operations.getChannel({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
});
return mapChannelToResponse({
return this.maskedChannelResponse(channel, params);
}
private async maskedChannelResponse(
channel: Channel,
params: {userId: UserID; viewer: ThreadViewer; requestCache: RequestCache},
): Promise<ChannelResponse> {
if (channel.isThread()) return this.threadResponse(channel, params.userId);
const response = await mapChannelToResponse({
channel,
currentUserId: params.userId,
userCacheService: this.userCacheService,
requestCache: params.requestCache,
});
const [masked] = await maskChannelResponseThreadBits(channel.guildId, params.viewer, [
await withThreadParentFields(
this.channelService.channelData.threadModify.repository.threads,
channel,
response,
params.viewer,
),
]);
return masked;
}
async getSlowmodeState(params: {user: User; channelId: ChannelID}): Promise<ChannelSlowmodeStateResponse> {
const state = await this.channelService.getSlowmodeState({user: params.user, channelId: params.channelId});
private async threadResponse(channel: Channel, userId: UserID): Promise<ChannelResponse> {
const repository = this.channelService.channelData.threadModify.repository;
const [state, member] = await Promise.all([
repository.threads.getState(channel.id),
repository.threads.getMember(channel.id, userId),
]);
if (!state) throw new UnknownChannelError();
const [view] = await loadThreadViews(repository, [state]);
if (!view) throw new UnknownChannelError();
return mapThreadToResponse({...view, channel}, member);
}
async getSlowmodeState(params: {
user: User;
viewer: ThreadViewer;
channelId: ChannelID;
}): Promise<ChannelSlowmodeStateResponse> {
const state = await this.channelService.getSlowmodeState({
user: params.user,
viewer: params.viewer,
channelId: params.channelId,
});
return {
rate_limit_per_user: state.rateLimitPerUser,
retry_after_ms: state.retryAfterMs,
@@ -44,9 +100,10 @@ export class ChannelRequestService {
};
}
async listRtcRegions(params: {userId: UserID; channelId: ChannelID}) {
async listRtcRegions(params: {userId: UserID; viewer: ThreadViewer; channelId: ChannelID}) {
const regions = await this.channelService.channelData.operations.getAvailableRtcRegions({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
});
return regions.map((region) => ({
@@ -58,8 +115,25 @@ export class ChannelRequestService {
async updateChannel(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
data: ChannelUpdateRequest;
data: ChannelUpdateGatedRequest;
clientFeatures: ReadonlySet<string>;
requestCache: RequestCache;
auditLogReason: string | null;
typeConversion?: ChannelTypeConversion | null;
}): Promise<ChannelResponse> {
if (!('available_tags' in params.data) || params.data.available_tags === undefined) {
return this.applyChannelUpdate(params);
}
return this.withForumTagsLock(params.channelId, () => this.applyChannelUpdate(params));
}
private async applyChannelUpdate(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
data: ChannelUpdateGatedRequest;
clientFeatures: ReadonlySet<string>;
requestCache: RequestCache;
auditLogReason: string | null;
@@ -67,6 +141,7 @@ export class ChannelRequestService {
}): Promise<ChannelResponse> {
const channel = await this.channelService.channelData.editChannel({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
data: params.data,
clientFeatures: params.clientFeatures,
@@ -74,16 +149,69 @@ export class ChannelRequestService {
auditLogReason: params.auditLogReason,
typeConversion: params.typeConversion,
});
return mapChannelToResponse({
channel,
currentUserId: params.userId,
userCacheService: this.userCacheService,
requestCache: params.requestCache,
return this.maskedChannelResponse(channel, params);
}
async editForumTags(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
edit: ForumTagEdit;
clientFeatures: ReadonlySet<string>;
requestCache: RequestCache;
auditLogReason: string | null;
}): Promise<ChannelResponse> {
const {channel, checkPermission} = await this.channelService.channelData.auth.getChannelAuthenticated({
userId: params.userId,
channelId: params.channelId,
viewer: params.viewer,
skipNsfwValidation: true,
});
if (!channel.isThreadOnly() || channel.guildId === null) throw new InvalidChannelTypeError();
await checkPermission(Permissions.MANAGE_CHANNELS);
const guildId = channel.guildId;
return this.withForumTagsLock(channel.id, async () => {
const config = await this.channelService.channelData.threadModify.repository.threads.getParentConfig(
guildId,
channel.id,
);
return this.applyChannelUpdate({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
data: {type: channel.type, available_tags: applyForumTagEdit(config, params.edit)} as ChannelUpdateGatedRequest,
clientFeatures: params.clientFeatures,
requestCache: params.requestCache,
auditLogReason: params.auditLogReason,
});
});
}
private async withForumTagsLock<T>(channelId: ChannelID, fn: () => Promise<T>): Promise<T> {
const lockKey = `channel:${channelId}:forum-tags`;
let lockToken: string | null = null;
for (let attempt = 0; attempt < FORUM_TAGS_LOCK_ACQUIRE_ATTEMPTS; attempt++) {
lockToken = await this.cacheService.acquireLock(lockKey, FORUM_TAGS_LOCK_TTL_SECONDS);
if (lockToken) break;
await new Promise((resolve) => setTimeout(resolve, FORUM_TAGS_LOCK_RETRY_DELAY_MS * (attempt + 1)));
}
if (!lockToken) {
throw new ThrottledError({
code: APIErrorCodes.RESOURCE_LOCKED,
retryAfterSeconds: 1,
data: {retry_after: 1},
});
}
try {
return await fn();
} finally {
await this.cacheService.releaseLock(lockKey, lockToken).catch(() => {});
}
}
async deleteChannel(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
requestCache: RequestCache;
silent?: boolean;
@@ -91,8 +219,19 @@ export class ChannelRequestService {
}): Promise<void> {
const channel = await this.channelService.channelData.operations.getChannel({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
});
if (channel.isThread()) {
await this.channelService.channelData.deleteThread({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
requestCache: params.requestCache,
auditLogReason: params.auditLogReason,
});
return;
}
if (channel.type === ChannelTypes.GROUP_DM) {
await this.channelService.groupDms.removeRecipientFromChannel({
userId: params.userId,
@@ -105,6 +244,7 @@ export class ChannelRequestService {
}
await this.channelService.channelData.operations.deleteChannel({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
requestCache: params.requestCache,
auditLogReason: params.auditLogReason,
@@ -5,6 +5,7 @@ import type {ChannelID} from '@app/api/BrandedTypes';
import type {IChannelRepository} from '@app/api/channel/IChannelRepository';
import type {AttachmentUploadTraceRepository} from '@app/api/channel/repositories/message/AttachmentUploadTraceRepository';
import {AttachmentUploadService} from '@app/api/channel/services/AttachmentUploadService';
import type {AuthenticatedChannel} from '@app/api/channel/services/AuthenticatedChannel';
import {CallService} from '@app/api/channel/services/CallService';
import {ChannelDataService} from '@app/api/channel/services/ChannelDataService';
import {GroupDmOperationsService} from '@app/api/channel/services/group_dm/GroupDmOperationsService';
@@ -12,6 +13,7 @@ import {MessageInteractionService} from '@app/api/channel/services/MessageIntera
import {MessageService} from '@app/api/channel/services/MessageService';
import {MessagePersistenceService} from '@app/api/channel/services/message/MessagePersistenceService';
import {UserMessageDeletionService} from '@app/api/channel/services/message/UserMessageDeletionService';
import {type ThreadViewer, viewerActive} from '@app/api/experiment/ChannelThreadsGate';
import type {IFavoriteMemeRepository} from '@app/api/favorite_meme/IFavoriteMemeRepository';
import type {GuildAuditLogService} from '@app/api/guild/GuildAuditLogService';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
@@ -30,9 +32,15 @@ import type {IUserRepository} from '@app/api/user/IUserRepository';
import type {VoiceAvailabilityService} from '@app/api/voice/VoiceAvailabilityService';
import type {IWebhookRepository} from '@app/api/webhook/IWebhookRepository';
import {Permissions} from '@fluxer/constants/src/ChannelConstants';
import {TEXT_THREAD_PARENT_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import type {IRateLimitService} from '@pkgs/rate_limit/src/IRateLimitService';
import type {IVirusScanService} from '@pkgs/virus_scan/src/IVirusScanService';
interface TypingCooldown {
message_send_cooldown_ms?: number;
thread_create_cooldown_ms?: number;
}
interface SlowmodeState {
rateLimitPerUser: number;
retryAfterMs: number;
@@ -188,8 +196,51 @@ export class ChannelService {
);
}
async getSlowmodeState({user, channelId}: {user: User; channelId: ChannelID}): Promise<SlowmodeState> {
const auth = await this.channelData.auth.getChannelAuthenticated({userId: user.id, channelId});
async getTypingCooldown({
user,
viewer,
auth,
}: {
user: User;
viewer: ThreadViewer;
auth: AuthenticatedChannel;
}): Promise<TypingCooldown | null> {
const {channel, guild} = auth;
if (!guild || user.isBot || !viewerActive(viewer, guild.id)) return null;
const createsThreads = TEXT_THREAD_PARENT_CHANNEL_TYPES.has(channel.type);
if (channel.rateLimitPerUser <= 0 || (await auth.hasPermission(Permissions.BYPASS_SLOWMODE))) return null;
const windowMs = channel.rateLimitPerUser * 1000;
const [messageSendCooldownMs, threadCreateCooldownMs] = await Promise.all([
this.peekSlowmode(`slowmode:${channel.id}:${user.id}`, windowMs),
createsThreads ? this.peekSlowmode(`slowmode-thread:${channel.id}:${user.id}`, windowMs) : Promise.resolve(0),
]);
if (messageSendCooldownMs <= 0 && threadCreateCooldownMs <= 0) return null;
return {
...(messageSendCooldownMs > 0 ? {message_send_cooldown_ms: messageSendCooldownMs} : {}),
...(threadCreateCooldownMs > 0 ? {thread_create_cooldown_ms: threadCreateCooldownMs} : {}),
};
}
private async peekSlowmode(identifier: string, windowMs: number): Promise<number> {
const peek = await this.rateLimitService.peekLimit({
identifier,
maxAttempts: 1,
windowMs,
algorithm: 'leaky_bucket',
});
return Math.max(0, peek.resetTime.getTime() - Date.now());
}
async getSlowmodeState({
user,
channelId,
viewer,
}: {
user: User;
viewer: ThreadViewer;
channelId: ChannelID;
}): Promise<SlowmodeState> {
const auth = await this.channelData.auth.getChannelAuthenticated({userId: user.id, channelId, viewer});
const rateLimitPerUser = auth.channel.rateLimitPerUser ?? 0;
if (!auth.guild || rateLimitPerUser <= 0 || user.isBot) {
return {rateLimitPerUser, retryAfterMs: 0, nextSendAllowedAt: null, canBypass: false};
@@ -198,8 +249,9 @@ export class ChannelService {
if (canBypass) {
return {rateLimitPerUser, retryAfterMs: 0, nextSendAllowedAt: null, canBypass: true};
}
const bucket = auth.channel.isThreadOnly() ? 'slowmode-thread' : 'slowmode';
const peek = await this.rateLimitService.peekLimit({
identifier: `slowmode:${channelId}:${user.id}`,
identifier: `${bucket}:${channelId}:${user.id}`,
maxAttempts: 1,
windowMs: rateLimitPerUser * 1000,
algorithm: 'leaky_bucket',
@@ -2,6 +2,7 @@
import type {ChannelID, MessageID, UserID} from '@app/api/BrandedTypes';
import type {IChannelRepository} from '@app/api/channel/IChannelRepository';
import type {AuthenticatedChannel} from '@app/api/channel/services/AuthenticatedChannel';
import {MessageInteractionAuthService} from '@app/api/channel/services/interaction/MessageInteractionAuthService';
import {MessagePinAuthService} from '@app/api/channel/services/interaction/MessagePinAuthService';
import {MessagePinService} from '@app/api/channel/services/interaction/MessagePinService';
@@ -9,6 +10,9 @@ import {MessageReactionService} from '@app/api/channel/services/interaction/Mess
import {MessageReadStateService} from '@app/api/channel/services/interaction/MessageReadStateService';
import {dispatchMessageUpdateBroadcast} from '@app/api/channel/services/message/MessageGatewayDispatch';
import type {MessagePersistenceService} from '@app/api/channel/services/message/MessagePersistenceService';
import {maskThreadArtifactsFor} from '@app/api/channel/services/message/ThreadMessageResponses';
import {assertThreadAllowed} from '@app/api/channel/services/thread/ThreadDenials';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import type {GuildAuditLogService} from '@app/api/guild/GuildAuditLogService';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
@@ -26,8 +30,10 @@ import {
} from '@app/api/user/NewConversationLimit';
import {assertGuildMemberCanCommunicate} from '@app/api/utils/GuildCommunicationUtils';
import {ChannelTypes, Permissions} from '@fluxer/constants/src/ChannelConstants';
import {threadWriteBlock} from '@fluxer/constants/src/ThreadPermissionUtils';
import {InvalidChannelTypeError} from '@fluxer/errors/src/domains/channel/InvalidChannelTypeError';
import {NewConversationsLimitedError} from '@fluxer/errors/src/domains/user/NewConversationsLimitedError';
import type {ChannelPinResponse} from '@fluxer/schema/src/domains/message/MessageResponseSchemas';
import type {ChannelPinResponse, MessageResponse} from '@fluxer/schema/src/domains/message/MessageResponseSchemas';
import type {UserPartialResponse} from '@fluxer/schema/src/domains/user/UserResponseSchemas';
export class MessageInteractionService {
@@ -71,22 +77,37 @@ export class MessageInteractionService {
);
}
async startTyping({userId, channelId}: {userId: UserID; channelId: ChannelID}): Promise<void> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId});
async startTyping({
userId,
channelId,
viewer,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
}): Promise<AuthenticatedChannel> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId, viewer});
if (authChannel.channel.isThreadOnly()) throw new InvalidChannelTypeError();
await authChannel.checkPermission(Permissions.SEND_MESSAGES);
assertGuildMemberCanCommunicate(authChannel.member);
if (!authChannel.guild && (await this.startsNewConversation(authChannel.channel, userId))) return;
if (authChannel.thread) {
assertThreadAllowed(threadWriteBlock('send', authChannel.thread.actor));
}
if (!authChannel.guild && (await this.startsNewConversation(authChannel.channel, userId))) return authChannel;
await this.readStateService.startTyping({authChannel, userId});
return authChannel;
}
async getChannelPins({
userId,
viewer,
channelId,
requestCache,
beforeTimestamp,
limit,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
requestCache: RequestCache;
beforeTimestamp?: Date;
@@ -95,24 +116,37 @@ export class MessageInteractionService {
items: Array<ChannelPinResponse>;
has_more: boolean;
}> {
const authChannel = await this.pinAuthService.getChannelAuthenticated({userId, channelId});
return this.pinService.getChannelPins({authChannel, userId, requestCache, beforeTimestamp, limit});
const authChannel = await this.pinAuthService.getChannelAuthenticated({userId, channelId, viewer});
const pins = await this.pinService.getChannelPins({authChannel, userId, requestCache, beforeTimestamp, limit});
const original = pins.items.map((item) => item.message as MessageResponse);
const messages = maskThreadArtifactsFor(viewer, authChannel.channel.guildId, original);
if (messages === original) return pins;
const byId = new Map(messages.map((message) => [message.id, message]));
return {
...pins,
items: pins.items.flatMap((item) => {
const message = byId.get(item.message.id);
return message ? [{...item, message}] : [];
}),
};
}
async pinMessage({
userId,
viewer,
channelId,
messageId,
requestCache,
auditLogReason,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
requestCache: RequestCache;
auditLogReason?: string | null;
}): Promise<void> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId, viewer});
if (!authChannel.guild && authChannel.channel.type !== ChannelTypes.DM_PERSONAL_NOTES) {
await this.authService.validateDMSendPermissions({channel: authChannel.channel, userId});
await this.assertConversationAllowed(authChannel.channel, userId);
@@ -122,18 +156,20 @@ export class MessageInteractionService {
async unpinMessage({
userId,
viewer,
channelId,
messageId,
requestCache,
auditLogReason,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
requestCache: RequestCache;
auditLogReason?: string | null;
}): Promise<void> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId, viewer});
if (!authChannel.guild && authChannel.channel.type !== ChannelTypes.DM_PERSONAL_NOTES) {
await this.authService.validateDMSendPermissions({channel: authChannel.channel, userId});
}
@@ -142,6 +178,7 @@ export class MessageInteractionService {
async getUsersForReaction({
userId,
viewer,
channelId,
messageId,
emoji,
@@ -149,6 +186,7 @@ export class MessageInteractionService {
after,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
emoji: string;
@@ -159,25 +197,27 @@ export class MessageInteractionService {
has_more: boolean;
next_after: string | null;
}> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId, viewer});
return this.reactionService.getUsersForReaction({authChannel, messageId, emoji, limit, after, userId});
}
async addReaction({
userId,
viewer,
sessionId,
channelId,
messageId,
emoji,
}: {
userId: UserID;
viewer: ThreadViewer;
sessionId?: string;
channelId: ChannelID;
messageId: MessageID;
emoji: string;
requestCache: RequestCache;
}): Promise<void> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId, viewer});
if (!authChannel.guild) {
await this.assertConversationAllowed(authChannel.channel, userId);
}
@@ -210,6 +250,7 @@ export class MessageInteractionService {
async removeReaction({
userId,
viewer,
sessionId,
channelId,
messageId,
@@ -217,6 +258,7 @@ export class MessageInteractionService {
targetId,
}: {
userId: UserID;
viewer: ThreadViewer;
sessionId?: string;
channelId: ChannelID;
messageId: MessageID;
@@ -224,12 +266,13 @@ export class MessageInteractionService {
targetId: UserID;
requestCache: RequestCache;
}): Promise<void> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId, viewer});
await this.reactionService.removeReaction({authChannel, messageId, emoji, targetId, sessionId, actorId: userId});
}
async removeOwnReaction({
userId,
viewer,
sessionId,
channelId,
messageId,
@@ -237,53 +280,60 @@ export class MessageInteractionService {
requestCache,
}: {
userId: UserID;
viewer: ThreadViewer;
sessionId?: string;
channelId: ChannelID;
messageId: MessageID;
emoji: string;
requestCache: RequestCache;
}): Promise<void> {
await this.removeReaction({userId, sessionId, channelId, messageId, emoji, targetId: userId, requestCache});
await this.removeReaction({userId, viewer, sessionId, channelId, messageId, emoji, targetId: userId, requestCache});
}
async removeAllReactionsForEmoji({
userId,
viewer,
channelId,
messageId,
emoji,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
emoji: string;
}): Promise<void> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId, viewer});
await this.reactionService.removeAllReactionsForEmoji({authChannel, messageId, emoji});
}
async removeAllReactions({
userId,
viewer,
channelId,
messageId,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
}): Promise<void> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId, viewer});
await this.reactionService.removeAllReactions({authChannel, messageId});
}
async getMessageReactions({
userId,
viewer,
channelId,
messageId,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
}): Promise<Array<MessageReaction>> {
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.authService.getChannelAuthenticated({userId, channelId, viewer});
return this.reactionService.getMessageReactions({authChannel, messageId});
}
@@ -20,6 +20,7 @@ import {MessageSendService} from '@app/api/channel/services/message/MessageSendS
import {MessageSystemService} from '@app/api/channel/services/message/MessageSystemService';
import {MessageValidationService} from '@app/api/channel/services/message/MessageValidationService';
import {MessageWriteLock} from '@app/api/channel/services/message/MessageWriteLock';
import {ThreadMessageActivity} from '@app/api/channel/services/message/ThreadMessageActivity';
import type {IFavoriteMemeRepository} from '@app/api/favorite_meme/IFavoriteMemeRepository';
import type {GuildAuditLogService} from '@app/api/guild/GuildAuditLogService';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
@@ -54,6 +55,7 @@ export class MessageService {
public readonly writeLock: MessageWriteLock;
public readonly crosspostPropagation: CrosspostPropagation;
public readonly crosspost: MessageCrosspostService;
public readonly threadActivity: ThreadMessageActivity;
constructor(
channelRepository: IChannelRepositoryAggregate,
@@ -117,7 +119,9 @@ export class MessageService {
snowflakeService,
favoriteMemeRepository,
});
this.threadActivity = new ThreadMessageActivity(channelRepository, gatewayService, userRepository);
this.send = new MessageSendService({
threadActivity: this.threadActivity,
channelRepository,
userRepository,
storageService,
@@ -4,6 +4,7 @@ import {createChannelID, createUserID} from '@app/api/BrandedTypes';
import type {ChannelService} from '@app/api/channel/services/ChannelService';
import type {StreamPreviewService} from '@app/api/channel/services/StreamPreviewService';
import {StreamService} from '@app/api/channel/services/StreamService';
import {SYSTEM_THREAD_VIEWER} from '@app/api/experiment/ChannelThreadsGate';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import {ChannelTypes} from '@fluxer/constants/src/ChannelConstants';
import {InvalidStreamThumbnailPayloadError} from '@fluxer/errors/src/domains/channel/InvalidStreamThumbnailPayloadError';
@@ -55,6 +56,7 @@ describe('StreamService.uploadPreview', () => {
const upload = (thumbnail: string) =>
streamService.uploadPreview({
viewer: SYSTEM_THREAD_VIEWER,
userId: USER_ID,
streamKey: STREAM_KEY,
channelId: CHANNEL_ID,
@@ -3,6 +3,7 @@
import {type ChannelID, createChannelID, createGuildID, type GuildID, type UserID} from '@app/api/BrandedTypes';
import type {ChannelService} from '@app/api/channel/services/ChannelService';
import type {StreamPreviewService} from '@app/api/channel/services/StreamPreviewService';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import {ChannelTypes, Permissions} from '@fluxer/constants/src/ChannelConstants';
import {InvalidChannelTypeError} from '@fluxer/errors/src/domains/channel/InvalidChannelTypeError';
@@ -77,11 +78,13 @@ export class StreamService {
private async assertStreamChannelAccess(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
parsedKey: ParsedStreamKey;
}): Promise<void> {
const channel = await this.channelService.channelData.operations.getChannel({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
});
if (channel.guildId) {
@@ -118,6 +121,7 @@ export class StreamService {
private async assertStreamMutationAccess(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
parsedKey: ParsedStreamKey;
}): Promise<void> {
@@ -149,11 +153,17 @@ export class StreamService {
}
}
async updateStreamRegion(params: {userId: UserID; streamKey: string; region?: string}): Promise<void> {
async updateStreamRegion(params: {
userId: UserID;
viewer: ThreadViewer;
streamKey: string;
region?: string;
}): Promise<void> {
const parsedKey = this.getParsedStreamKeyOrThrow(params.streamKey);
const channelId = this.getChannelIdFromParsedKeyOrThrow(parsedKey);
await this.assertStreamMutationAccess({
userId: params.userId,
viewer: params.viewer,
channelId,
parsedKey,
});
@@ -164,7 +174,7 @@ export class StreamService {
);
}
async getPreview(params: {userId: UserID; streamKey: string}): Promise<{
async getPreview(params: {userId: UserID; viewer: ThreadViewer; streamKey: string}): Promise<{
buffer: Uint8Array;
contentType: string;
} | null> {
@@ -172,6 +182,7 @@ export class StreamService {
const channelId = this.getChannelIdFromParsedKeyOrThrow(parsedKey);
await this.assertStreamChannelAccess({
userId: params.userId,
viewer: params.viewer,
channelId,
parsedKey,
});
@@ -180,6 +191,7 @@ export class StreamService {
async uploadPreview(params: {
userId: UserID;
viewer: ThreadViewer;
streamKey: string;
channelId: ChannelID;
thumbnail: string;
@@ -188,6 +200,7 @@ export class StreamService {
const parsedKey = this.getParsedStreamKeyOrThrow(params.streamKey);
await this.assertStreamMutationAccess({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
parsedKey,
});
@@ -206,6 +219,7 @@ export class StreamService {
async createPreviewUploadUrl(params: {
userId: UserID;
viewer: ThreadViewer;
streamKey: string;
channelId: ChannelID;
contentType?: string;
@@ -214,6 +228,7 @@ export class StreamService {
const parsedKey = this.getParsedStreamKeyOrThrow(params.streamKey);
await this.assertStreamMutationAccess({
userId: params.userId,
viewer: params.viewer,
channelId: params.channelId,
parsedKey,
});
@@ -226,11 +241,12 @@ export class StreamService {
});
}
async deletePreview(params: {userId: UserID; streamKey: string}): Promise<void> {
async deletePreview(params: {userId: UserID; viewer: ThreadViewer; streamKey: string}): Promise<void> {
const parsedKey = this.getParsedStreamKeyOrThrow(params.streamKey);
const channelId = this.getChannelIdFromParsedKeyOrThrow(parsedKey);
await this.assertStreamMutationAccess({
userId: params.userId,
viewer: params.viewer,
channelId,
parsedKey,
});
@@ -10,14 +10,38 @@ import {
} from '@app/api/channel/services/ChannelFollowers';
import type {ChannelAuthService} from '@app/api/channel/services/channel_data/ChannelAuthService';
import type {ChannelUtilsService} from '@app/api/channel/services/channel_data/ChannelUtilsService';
import {dispatchThreadEvents} from '@app/api/channel/services/thread/ThreadDispatch';
import {
loadConvertibleParentThreads,
PARENT_CONVERSION_LOCK_TTL_SECONDS,
retypedThreadEvents,
retypeParentThreads,
} from '@app/api/channel/services/thread/ThreadParentConversion';
import {
buildThreadParentPatch,
loadThreadParentConfig,
serializeThreadParentForAudit,
type ThreadParentSettingsInput,
} from '@app/api/channel/services/thread/ThreadParentSettings';
import {enqueueDeleteChannelThreads} from '@app/api/channel/threads/ThreadJobs';
import {
everEnabled,
guildActive,
isTainted,
THREAD_FEATURE_CHANNEL_TYPES,
type ThreadViewer,
viewerActive,
} from '@app/api/experiment/ChannelThreadsGate';
import type {GuildAuditLogService} from '@app/api/guild/GuildAuditLogService';
import {mapGuildToGuildResponse} from '@app/api/guild/GuildModel';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
import {ChannelHelpers} from '@app/api/guild/services/channel/ChannelHelpers';
import {createGuildMfaEnforcer} from '@app/api/guild/services/GuildMfaEnforcement';
import {hasThreadPermissionBits, resolveProtectedBitActor} from '@app/api/guild/services/ThreadPermissionBits';
import {contentModerationService} from '@app/api/infrastructure/ContentModerationService';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import type {ILiveKitService} from '@app/api/infrastructure/ILiveKitService';
import type {ISnowflakeService} from '@app/api/infrastructure/ISnowflakeService';
import type {IVoiceRoomStore} from '@app/api/infrastructure/IVoiceRoomStore';
import type {IInviteRepository} from '@app/api/invite/IInviteRepository';
import {Logger} from '@app/api/Logger';
@@ -26,17 +50,17 @@ import {createLimitMatchContext} from '@app/api/limits/LimitMatchContextBuilder'
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
import type {Channel} from '@app/api/models/Channel';
import {ChannelPermissionOverwrite} from '@app/api/models/ChannelPermissionOverwrite';
import type {ThreadState} from '@app/api/models/ThreadState';
import {deleteChannelMessageSearchDocuments} from '@app/api/search/MessageSearchIndexCleanup';
import type {IUserRepository} from '@app/api/user/IUserRepository';
import {serializeChannelForAudit} from '@app/api/utils/AuditSerializationUtils';
import {applyProtectedOverwriteBits} from '@app/api/utils/featureUtils';
import {applyProtectedOverwriteBits, permissionWriteMask, protectedThreadBits} from '@app/api/utils/featureUtils';
import {overwriteGrantedBits} from '@app/api/utils/PermissionUtils';
import type {VoiceAvailabilityService} from '@app/api/voice/VoiceAvailabilityService';
import type {VoiceRegionAvailability} from '@app/api/voice/VoiceModel';
import type {IWebhookRepository} from '@app/api/webhook/IWebhookRepository';
import {AuditLogActionType} from '@fluxer/constants/src/AuditLogActionType';
import {
ALL_PERMISSIONS,
ANNOUNCEMENT_CONVERTIBLE_CHANNEL_TYPES,
ChannelTypes,
GUILD_TEXT_BASED_CHANNEL_TYPES,
@@ -45,6 +69,8 @@ import {
} from '@fluxer/constants/src/ChannelConstants';
import {ContentWarningLevel, clampVoiceChannelBitrate, GuildFeatures} from '@fluxer/constants/src/GuildConstants';
import {MAX_CHANNELS_PER_CATEGORY} from '@fluxer/constants/src/LimitConstants';
import {THREAD_ONLY_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {withImplicitThreadBits} from '@fluxer/constants/src/ThreadPermissionUtils';
import {ValidationErrorCodes} from '@fluxer/constants/src/ValidationErrorCodes';
import {ChannelHasFollowedChannelsError} from '@fluxer/errors/src/domains/channel/ChannelHasFollowedChannelsError';
import {ChannelTypeConversionNotSupportedError} from '@fluxer/errors/src/domains/channel/ChannelTypeConversionNotSupportedError';
@@ -84,6 +110,13 @@ export interface ChannelUpdateData {
nicks?: Record<string, string | null> | null;
}
function assertOverwriteTarget(channel: Channel, viewer: ThreadViewer | undefined): void {
if (!THREAD_FEATURE_CHANNEL_TYPES.has(channel.type)) return;
const active = viewer !== undefined && channel.guildId !== null && viewerActive(viewer, channel.guildId);
if (!active) throw new UnknownChannelError();
if (channel.isThread()) throw new InvalidChannelTypeError();
}
export class ChannelOperationsService {
constructor(
private channelRepository: IChannelRepositoryAggregate,
@@ -101,20 +134,24 @@ export class ChannelOperationsService {
private limitConfigService: LimitConfigService,
private rateLimitService: IRateLimitService,
private cacheService: ICacheService,
private snowflakeService: ISnowflakeService,
) {}
async getChannel({
userId,
viewer,
channelId,
skipNsfwValidation,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
skipNsfwValidation?: boolean;
}): Promise<Channel> {
const {channel} = await this.channelAuthService.getChannelAuthenticated({
userId,
channelId,
viewer,
skipNsfwValidation,
});
return channel;
@@ -138,24 +175,29 @@ export class ChannelOperationsService {
async editChannel({
userId,
viewer,
channelId,
data,
clientFeatures,
requestCache,
auditLogReason,
typeConversion,
threadParent,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
data: ChannelUpdateData;
clientFeatures: ReadonlySet<string>;
requestCache: RequestCache;
auditLogReason: string | null;
typeConversion?: ChannelTypeConversion | null;
threadParent?: ThreadParentSettingsInput | null;
}): Promise<Channel> {
const {channel, guild, checkPermission} = await this.channelAuthService.getChannelAuthenticated({
userId,
channelId,
viewer,
skipNsfwValidation: true,
});
if (channel.type === ChannelTypes.GROUP_DM) {
@@ -165,6 +207,18 @@ export class ChannelOperationsService {
await checkPermission(Permissions.MANAGE_CHANNELS);
const nextType = resolveNextChannelType(channel, typeConversion ?? null);
const guildIdValue = createGuildID(BigInt(guild.id));
const parentConfig = await loadThreadParentConfig(this.channelRepository.threads, channel);
const parentPatch =
threadParent && guildActive(guildIdValue)
? await buildThreadParentPatch({
channelType: channel.type,
guildId: guildIdValue,
input: threadParent,
current: parentConfig,
guildRepository: this.guildRepository,
generateId: () => this.snowflakeService.generate(),
})
: null;
contentModerationService.scanText(data.name ?? null, {
userId,
guildId: guildIdValue,
@@ -180,7 +234,7 @@ export class ChannelOperationsService {
surface: 'profile_field',
});
let channelName = data.name ?? channel.name;
if (data.name !== undefined && isTextNamedChannelType(channel.type)) {
if (data.name !== undefined && (isTextNamedChannelType(channel.type) || channel.isThreadOnly())) {
const hasFlexibleNamesEnabled = guild.features?.includes(GuildFeatures.TEXT_CHANNEL_FLEXIBLE_NAMES) ?? false;
if (!hasFlexibleNamesEnabled) {
channelName = ChannelNameType.parse(data.name);
@@ -219,25 +273,34 @@ export class ChannelOperationsService {
const guildId = createGuildID(BigInt(guild.id));
await checkPermission(Permissions.MANAGE_ROLES);
const isOwner = guild.owner_id === userId.toString();
const channelPermissions = await this.gatewayService.getUserPermissions({
const gatewayPermissions = await this.gatewayService.getUserPermissions({
guildId,
userId,
channelId: channel.id,
});
const actor = await resolveProtectedBitActor({
guildId,
userId,
clientFeatures,
viewer,
isBot: async () => viewer.kind === 'user' && viewer.bot,
});
const writeMask = permissionWriteMask(actor);
const channelPermissions = actor.threadBits ? withImplicitThreadBits(gatewayPermissions) : gatewayPermissions;
permissionOverwrites = new Map();
for (const overwrite of data.permission_overwrites ?? []) {
const targetId = overwrite.type === 0 ? createRoleID(overwrite.id) : createUserID(overwrite.id);
const existing = previousPermissionOverwrites?.get(targetId);
const protectedBits = applyProtectedOverwriteBits(
{
allow: (overwrite.allow ? BigInt(overwrite.allow) : 0n) & ALL_PERMISSIONS,
deny: (overwrite.deny ? BigInt(overwrite.deny) : 0n) & ALL_PERMISSIONS,
allow: (overwrite.allow ? BigInt(overwrite.allow) : 0n) & writeMask,
deny: (overwrite.deny ? BigInt(overwrite.deny) : 0n) & writeMask,
},
{
allow: existing?.allow ?? 0n,
deny: existing?.deny ?? 0n,
},
clientFeatures,
actor,
);
permissionOverwrites.set(
targetId,
@@ -248,6 +311,19 @@ export class ChannelOperationsService {
}),
);
}
const keptBits = protectedThreadBits(actor);
if (keptBits !== 0n) {
for (const [targetId, previous] of previousPermissionOverwrites ?? []) {
if (permissionOverwrites.has(targetId)) continue;
const allow = previous.allow & keptBits;
const deny = previous.deny & keptBits;
if (allow === 0n && deny === 0n) continue;
permissionOverwrites.set(
targetId,
new ChannelPermissionOverwrite({type: previous.type, allow_: allow, deny_: deny}),
);
}
}
if (!isOwner) {
const targetIds = new Set([...(previousPermissionOverwrites?.keys() ?? []), ...permissionOverwrites.keys()]);
for (const targetId of targetIds) {
@@ -255,7 +331,7 @@ export class ChannelOperationsService {
previousPermissionOverwrites?.get(targetId),
permissionOverwrites.get(targetId),
);
if ((grantedBits & ~channelPermissions) !== 0n) {
if ((grantedBits & ~keptBits & ~channelPermissions) !== 0n) {
throw new MissingPermissionsError();
}
}
@@ -292,7 +368,7 @@ export class ChannelOperationsService {
? data.voice_connection_limit
: channel.voiceConnectionLimit,
rate_limit_per_user:
data.rate_limit_per_user !== undefined && GUILD_TEXT_BASED_CHANNEL_TYPES.has(channel.type)
data.rate_limit_per_user !== undefined && acceptsRateLimit(channel.type)
? data.rate_limit_per_user
: channel.rateLimitPerUser,
nsfw: resolveNsfwOverrideWrite(channel, data),
@@ -309,26 +385,54 @@ export class ChannelOperationsService {
]),
),
};
let retypedThreads: Array<ThreadState> = [];
const toAnnouncement =
nextType === ChannelTypes.GUILD_ANNOUNCEMENT && channel.type !== ChannelTypes.GUILD_ANNOUNCEMENT;
const updatedChannel =
nextType === ChannelTypes.GUILD_ANNOUNCEMENT && channel.type !== ChannelTypes.GUILD_ANNOUNCEMENT
? await withChannelFollowLock(this.cacheService, channelId, async () => {
const webhooks = await this.webhookRepository.listByChannel(channelId);
if (webhooks.some((webhook) => webhook.type === WebhookTypes.CHANNEL_FOLLOWER)) {
throw new ChannelHasFollowedChannelsError();
}
return await this.channelRepository.channelData.upsert(updatedChannelData);
})
toAnnouncement || (nextType !== channel.type && everEnabled())
? await withChannelFollowLock(
this.cacheService,
channelId,
async () => {
if (toAnnouncement) {
const webhooks = await this.webhookRepository.listByChannel(channelId);
if (webhooks.some((webhook) => webhook.type === WebhookTypes.CHANNEL_FOLLOWER)) {
throw new ChannelHasFollowedChannelsError();
}
}
const threads = await loadConvertibleParentThreads(
this.channelRepository,
guildIdValue,
channel,
nextType,
);
if (threads.length === 0) return this.channelRepository.channelData.upsert(updatedChannelData);
const {parent, active} = await retypeParentThreads(this.channelRepository, threads, nextType, () =>
this.channelRepository.channelData.upsert(updatedChannelData),
);
retypedThreads = active;
return parent;
},
everEnabled() ? PARENT_CONVERSION_LOCK_TTL_SECONDS : undefined,
)
: await this.channelRepository.channelData.upsert(updatedChannelData);
if (channel.type === ChannelTypes.GUILD_ANNOUNCEMENT && nextType !== ChannelTypes.GUILD_ANNOUNCEMENT) {
await enqueueChannelFollowerRemoval({sourceChannelId: channelId, reason: 'converted'});
}
if (parentPatch) await this.channelRepository.threads.patchParentConfig(guildIdValue, channel.id, parentPatch);
const nextParentConfig = parentPatch
? await loadThreadParentConfig(this.channelRepository.threads, updatedChannel)
: parentConfig;
if (
data.rate_limit_per_user !== undefined &&
GUILD_TEXT_BASED_CHANNEL_TYPES.has(channel.type) &&
acceptsRateLimit(channel.type) &&
data.rate_limit_per_user !== channel.rateLimitPerUser
) {
try {
await this.rateLimitService.clearLimitsByIdentifierPrefix(`slowmode:${channelId}:`);
if (everEnabled() && (await isTainted(guildIdValue))) {
await this.rateLimitService.clearLimitsByIdentifierPrefix(`slowmode-thread:${channelId}:`);
}
} catch (error) {
Logger.error(
{error, channelId: channelId.toString()},
@@ -337,6 +441,11 @@ export class ChannelOperationsService {
}
}
await this.channelUtilsService.dispatchChannelUpdate({channel: updatedChannel, requestCache});
await dispatchThreadEvents(
this.gatewayService,
guildIdValue,
await retypedThreadEvents(this.channelRepository, updatedChannel, retypedThreads),
);
if (channel.type === ChannelTypes.GUILD_CATEGORY && data.permission_overwrites !== undefined && guild) {
await this.propagatePermissionsToSyncedChildren({
categoryChannel: updatedChannel,
@@ -356,8 +465,14 @@ export class ChannelOperationsService {
channelId,
});
}
const beforeSnapshot = serializeChannelForAudit(channel);
const afterSnapshot = serializeChannelForAudit(updatedChannel);
const beforeSnapshot = {
...serializeChannelForAudit(channel),
...serializeThreadParentForAudit(channel.type, parentConfig),
};
const afterSnapshot = {
...serializeChannelForAudit(updatedChannel),
...serializeThreadParentForAudit(updatedChannel.type, nextParentConfig),
};
const changes = this.guildAuditLogService.computeChanges(beforeSnapshot, afterSnapshot);
if (changes.length > 0) {
const builder = this.guildAuditLogService
@@ -398,11 +513,13 @@ export class ChannelOperationsService {
async deleteChannel({
userId,
viewer,
channelId,
requestCache,
auditLogReason,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
requestCache: RequestCache;
auditLogReason: string | null;
@@ -410,6 +527,7 @@ export class ChannelOperationsService {
const {channel, guild, checkPermission} = await this.channelAuthService.getChannelAuthenticated({
userId,
channelId,
viewer,
skipNsfwValidation: true,
});
if (this.channelAuthService.isPersonalNotesChannel({userId, channelId})) {
@@ -419,13 +537,14 @@ export class ChannelOperationsService {
await checkPermission(Permissions.MANAGE_CHANNELS);
const guildId = createGuildID(BigInt(guild.id));
if (channel.type === ChannelTypes.GUILD_CATEGORY) {
const guildChannels = await this.channelRepository.channelData.listGuildChannels(guildId);
const guildChannels = await this.channelRepository.channelData.listGuildChannels(guildId, 'maintenance');
const childChannels = guildChannels.filter((ch: Channel) => ch.parentId === channelId);
for (const childChannel of childChannels) {
const updatedChild = await this.channelRepository.channelData.upsert({
...childChannel.toRow(),
parent_id: null,
});
if (THREAD_ONLY_CHANNEL_TYPES.has(updatedChild.type) && !guildActive(guildId)) continue;
await this.channelUtilsService.dispatchChannelUpdate({channel: updatedChild, requestCache});
}
}
@@ -470,7 +589,11 @@ export class ChannelOperationsService {
'Failed to record guild audit log',
);
}
await this.channelRepository.channelData.delete(channelId, guildId);
await this.channelRepository.channelData.delete(channelId, guildId, channel.type);
if (channel.isThreadParent() && everEnabled() && (await isTainted(guildId, {fresh: true}))) {
await enqueueDeleteChannelThreads(guildId, channelId);
await this.channelRepository.threads.deleteParentConfig(guildId, channelId);
}
const guildModel = await this.guildRepository.findUnique(guildId);
if (guildModel) {
const guildRow = guildModel.toRow();
@@ -495,9 +618,11 @@ export class ChannelOperationsService {
async getAvailableRtcRegions({
userId,
viewer,
channelId,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
}): Promise<Array<VoiceRegionAvailability>> {
if (this.voiceAvailabilityService === null) {
@@ -506,6 +631,7 @@ export class ChannelOperationsService {
const {channel, guild} = await this.channelAuthService.getChannelAuthenticated({
userId,
channelId,
viewer,
skipNsfwValidation: true,
});
if (channel.type !== ChannelTypes.GUILD_VOICE) {
@@ -535,7 +661,7 @@ export class ChannelOperationsService {
guildId: GuildID;
requestCache: RequestCache;
}): Promise<void> {
const guildChannels = await this.channelRepository.channelData.listGuildChannels(guildId);
const guildChannels = await this.channelRepository.channelData.listGuildChannels(guildId, 'maintenance');
const childChannels = guildChannels.filter((ch: Channel) => ch.parentId === categoryChannel.id);
const syncedChannels: Array<Channel> = [];
for (const child of childChannels) {
@@ -555,6 +681,7 @@ export class ChannelOperationsService {
]),
),
});
if (THREAD_ONLY_CHANNEL_TYPES.has(updatedChild.type) && !guildActive(guildId)) return;
await this.channelUtilsService.dispatchChannelUpdate({channel: updatedChild, requestCache});
}),
);
@@ -617,7 +744,7 @@ export class ChannelOperationsService {
if (params.channel.type === ChannelTypes.GUILD_CATEGORY) {
throw InputValidationError.fromCode('parent_id', ValidationErrorCodes.CATEGORIES_CANNOT_HAVE_PARENTS);
}
const guildChannels = await this.channelRepository.channelData.listGuildChannels(params.guildId);
const guildChannels = await this.channelRepository.channelData.listGuildChannels(params.guildId, 'enrolled');
const parentChannel = guildChannels.find((channel) => channel.id === params.parentId);
if (!parentChannel) {
throw InputValidationError.fromCode('parent_id', ValidationErrorCodes.INVALID_PARENT_CHANNEL);
@@ -661,11 +788,13 @@ export class ChannelOperationsService {
deny_: bigint;
};
clientFeatures: ReadonlySet<string>;
viewer?: ThreadViewer;
requestCache: RequestCache;
auditLogReason: string | null;
}): Promise<void> {
const channel = await this.channelRepository.channelData.findUnique(params.channelId);
if (!channel?.guildId) throw new UnknownChannelError();
assertOverwriteTarget(channel, params.viewer);
await this.checkOverwritePermission({guildId: channel.guildId, userId: params.userId, channelId: channel.id});
const userPermissions = await this.gatewayService.getUserPermissions({
guildId: channel.guildId,
@@ -674,22 +803,31 @@ export class ChannelOperationsService {
});
const targetId = params.overwrite.type === 0 ? createRoleID(params.overwriteId) : createUserID(params.overwriteId);
const existing = channel.permissionOverwrites?.get(targetId);
const actor = await resolveProtectedBitActor({
guildId: channel.guildId,
userId: params.userId,
clientFeatures: params.clientFeatures,
viewer: params.viewer,
isBot: async () => (await this.userRepository.findUnique(params.userId))?.isBot ?? false,
});
const writeMask = permissionWriteMask(actor);
const protectedBits = applyProtectedOverwriteBits(
{
allow: params.overwrite.allow_ & ALL_PERMISSIONS,
deny: params.overwrite.deny_ & ALL_PERMISSIONS,
allow: params.overwrite.allow_ & writeMask,
deny: params.overwrite.deny_ & writeMask,
},
{
allow: existing?.allow ?? 0n,
deny: existing?.deny ?? 0n,
},
params.clientFeatures,
actor,
);
const sanitizedAllow = protectedBits.allow;
const sanitizedDeny = protectedBits.deny;
const hasAdministrator = (userPermissions & Permissions.ADMINISTRATOR) !== 0n;
const grantedBits = overwriteGrantedBits(existing, {allow: sanitizedAllow, deny: sanitizedDeny});
if (!hasAdministrator && (grantedBits & ~userPermissions) !== 0n) throw new MissingPermissionsError();
const effectivePermissions = actor.threadBits ? withImplicitThreadBits(userPermissions) : userPermissions;
if (!hasAdministrator && (grantedBits & ~effectivePermissions) !== 0n) throw new MissingPermissionsError();
const previousPermissionOverwrites = channel.permissionOverwrites;
const nextOverwrite = new ChannelPermissionOverwrite({
type: params.overwrite.type,
@@ -698,12 +836,15 @@ export class ChannelOperationsService {
});
const overwrites = new Map(channel.permissionOverwrites ?? []);
overwrites.set(targetId, nextOverwrite);
const updated = await this.channelRepository.channelData.upsert({
...channel.toRow(),
permission_overwrites: new Map(
Array.from(overwrites.entries()).map(([id, ow]) => [id, ow.toPermissionOverwrite()]),
),
});
const updated = await this.channelRepository.channelData.upsert(
{
...channel.toRow(),
permission_overwrites: new Map(
Array.from(overwrites.entries()).map(([id, ow]) => [id, ow.toPermissionOverwrite()]),
),
},
channel.toRow(),
);
await this.channelUtilsService.dispatchChannelUpdate({channel: updated, requestCache: params.requestCache});
if (channel.type === ChannelTypes.GUILD_CATEGORY) {
await this.propagatePermissionsToSyncedChildren({
@@ -727,17 +868,32 @@ export class ChannelOperationsService {
userId: UserID;
channelId: ChannelID;
overwriteId: bigint;
clientFeatures?: ReadonlySet<string>;
viewer?: ThreadViewer;
requestCache: RequestCache;
auditLogReason: string | null;
}): Promise<void> {
const channel = await this.channelRepository.channelData.findUnique(params.channelId);
if (!channel?.guildId) throw new UnknownChannelError();
assertOverwriteTarget(channel, params.viewer);
await this.checkOverwritePermission({guildId: channel.guildId, userId: params.userId, channelId: channel.id});
const previousPermissionOverwrites = channel.permissionOverwrites;
const overwrites = new Map(channel.permissionOverwrites ?? []);
const removedRole = overwrites.get(createRoleID(params.overwriteId));
const removedUser = overwrites.get(createUserID(params.overwriteId));
const removed = removedRole ?? removedUser;
const kept =
removed && (hasThreadPermissionBits(removed.allow) || hasThreadPermissionBits(removed.deny))
? protectedThreadBits(
await resolveProtectedBitActor({
guildId: channel.guildId,
userId: params.userId,
clientFeatures: params.clientFeatures ?? new Set(),
viewer: params.viewer,
isBot: async () => (await this.userRepository.findUnique(params.userId))?.isBot ?? false,
}),
)
: 0n;
if (removed) {
const userPermissions = await this.gatewayService.getUserPermissions({
guildId: channel.guildId,
@@ -745,16 +901,30 @@ export class ChannelOperationsService {
channelId: channel.id,
});
const hasAdministrator = (userPermissions & Permissions.ADMINISTRATOR) !== 0n;
if (!hasAdministrator && (removed.deny & ~userPermissions) !== 0n) throw new MissingPermissionsError();
if (!hasAdministrator && (removed.deny & ~kept & ~userPermissions) !== 0n) throw new MissingPermissionsError();
}
overwrites.delete(createRoleID(params.overwriteId));
overwrites.delete(createUserID(params.overwriteId));
const updated = await this.channelRepository.channelData.upsert({
...channel.toRow(),
permission_overwrites: new Map(
Array.from(overwrites.entries()).map(([id, ow]) => [id, ow.toPermissionOverwrite()]),
),
});
if (removed && kept !== 0n) {
const removedTargetId = removed.type === 0 ? createRoleID(params.overwriteId) : createUserID(params.overwriteId);
overwrites.set(
removedTargetId,
new ChannelPermissionOverwrite({
type: removed.type,
allow_: removed.allow & kept,
deny_: removed.deny & kept,
}),
);
}
const updated = await this.channelRepository.channelData.upsert(
{
...channel.toRow(),
permission_overwrites: new Map(
Array.from(overwrites.entries()).map(([id, ow]) => [id, ow.toPermissionOverwrite()]),
),
},
channel.toRow(),
);
await this.channelUtilsService.dispatchChannelUpdate({channel: updated, requestCache: params.requestCache});
if (channel.type === ChannelTypes.GUILD_CATEGORY) {
await this.propagatePermissionsToSyncedChildren({
@@ -807,10 +977,15 @@ function isWritableGuildChannel(type: number): boolean {
type === ChannelTypes.GUILD_ANNOUNCEMENT ||
type === ChannelTypes.GUILD_VOICE ||
type === ChannelTypes.GUILD_LINK ||
type === ChannelTypes.GUILD_CATEGORY
type === ChannelTypes.GUILD_CATEGORY ||
THREAD_ONLY_CHANNEL_TYPES.has(type)
);
}
function acceptsRateLimit(type: number): boolean {
return GUILD_TEXT_BASED_CHANNEL_TYPES.has(type) || THREAD_ONLY_CHANNEL_TYPES.has(type);
}
function resolveNsfwOverrideWrite(channel: Channel, data: ChannelUpdateData): boolean | null {
if (!isWritableGuildChannel(channel.type)) {
return channel.nsfwOverride;
@@ -6,6 +6,7 @@ import type {IChannelRepositoryAggregate} from '@app/api/channel/repositories/IC
import {dispatchChannelEvent} from '@app/api/channel/services/ChannelGatewayDispatch';
import {dispatchMessageCreateBroadcast} from '@app/api/channel/services/message/MessageGatewayDispatch';
import {purgeMessageAttachments} from '@app/api/channel/services/message/MessageHelpers';
import {withThreadParentFields} from '@app/api/channel/services/thread/ThreadParentSettings';
import type {IPurgeQueue} from '@app/api/infrastructure/CachePurgeQueue';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import type {IStorageService} from '@app/api/infrastructure/IStorageService';
@@ -48,12 +49,16 @@ export class ChannelUtilsService {
async dispatchChannelUpdate({channel, requestCache}: {channel: Channel; requestCache: RequestCache}): Promise<void> {
if (channel.guildId) {
const channelResponse = await mapChannelToResponse({
const channelResponse = await withThreadParentFields(
this.channelRepository.threads,
channel,
currentUserId: null,
userCacheService: this.userCacheService,
requestCache,
});
await mapChannelToResponse({
channel,
currentUserId: null,
userCacheService: this.userCacheService,
requestCache,
}),
);
await dispatchChannelEvent({
gatewayService: this.gatewayService,
channel,
@@ -3,6 +3,7 @@
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import type {Channel} from '@app/api/models/Channel';
import {TEXT_BASED_CHANNEL_TYPES} from '@fluxer/constants/src/ChannelConstants';
import {THREAD_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {CannotSendMessageToNonTextChannelError} from '@fluxer/errors/src/domains/channel/CannotSendMessageToNonTextChannelError';
import type {GuildResponse} from '@fluxer/schema/src/domains/guild/GuildResponseSchemas';
@@ -21,7 +22,7 @@ export abstract class MessageInteractionBase {
}
protected ensureTextChannel(channel: Channel): void {
if (!TEXT_BASED_CHANNEL_TYPES.has(channel.type)) {
if (!TEXT_BASED_CHANNEL_TYPES.has(channel.type) && !THREAD_CHANNEL_TYPES.has(channel.type)) {
throw new CannotSendMessageToNonTextChannelError();
}
}
@@ -12,6 +12,7 @@ import {
} from '@app/api/channel/services/message/MessageGatewayDispatch';
import type {MessagePersistenceService} from '@app/api/channel/services/message/MessagePersistenceService';
import {createMessageResponseDataService} from '@app/api/channel/services/message/MessageResponseDataService';
import {assertThreadInteractionAllowed} from '@app/api/channel/services/thread/ThreadInteractionGuards';
import type {GuildAuditLogService} from '@app/api/guild/GuildAuditLogService';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import type {ISnowflakeService} from '@app/api/infrastructure/ISnowflakeService';
@@ -21,7 +22,9 @@ import type {Message} from '@app/api/models/Message';
import {AuditLogActionType} from '@fluxer/constants/src/AuditLogActionType';
import {MessageTypes, Permissions} from '@fluxer/constants/src/ChannelConstants';
import {GuildOperations} from '@fluxer/constants/src/GuildConstants';
import {THREAD_ONLY_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {CannotEditSystemMessageError} from '@fluxer/errors/src/domains/channel/CannotEditSystemMessageError';
import {InvalidChannelTypeError} from '@fluxer/errors/src/domains/channel/InvalidChannelTypeError';
import {UnknownMessageError} from '@fluxer/errors/src/domains/channel/UnknownMessageError';
import {FeatureTemporarilyDisabledError} from '@fluxer/errors/src/domains/core/FeatureTemporarilyDisabledError';
import type {ChannelPinResponse} from '@fluxer/schema/src/domains/message/MessageResponseSchemas';
@@ -75,6 +78,7 @@ export class MessagePinService extends MessageInteractionBase {
has_more: boolean;
}> {
const {channel} = authChannel;
if (THREAD_ONLY_CHANNEL_TYPES.has(channel.type)) throw new InvalidChannelTypeError();
this.ensureTextChannel(channel);
const hasReadHistory = !authChannel.guild || (await authChannel.hasPermission(Permissions.READ_MESSAGE_HISTORY));
if (!hasReadHistory) {
@@ -158,6 +162,7 @@ export class MessagePinService extends MessageInteractionBase {
}
}
this.ensureTextChannel(channel);
assertThreadInteractionAllowed(authChannel, 'pin');
await this.assertMessageHistoryAccess({authChannel, messageId});
const message = await this.channelRepository.messages.getMessage(channel.id, messageId);
if (!message) throw new UnknownMessageError();
@@ -168,7 +173,7 @@ export class MessagePinService extends MessageInteractionBase {
const updatedMessage = await this.channelRepository.messages.upsertMessage(updatedMessageData, message.toRow());
await this.channelRepository.messageInteractions.addChannelPin(channel.id, messageId, now);
const updatedChannelData = {...channel.toRow(), last_pin_timestamp: now};
const updatedChannel = await this.channelRepository.channelData.upsert(updatedChannelData);
const updatedChannel = await this.channelRepository.channelData.upsert(updatedChannelData, channel.toRow());
await this.dispatchChannelPinsUpdate(updatedChannel);
await this.sendPinSystemMessage({channel, message, userId});
await dispatchMessageUpdateBroadcast({
@@ -184,6 +189,7 @@ export class MessagePinService extends MessageInteractionBase {
channel_id: channel.id.toString(),
message_id: messageId.toString(),
})
.withThreadScope(channel.isThread())
.withReason(auditLogReason ?? null)
.commit();
}
@@ -209,6 +215,7 @@ export class MessagePinService extends MessageInteractionBase {
}
}
this.ensureTextChannel(channel);
assertThreadInteractionAllowed(authChannel, 'pin');
await this.assertMessageHistoryAccess({authChannel, messageId});
const message = await this.channelRepository.messages.getMessage(channel.id, messageId);
if (!message) throw new UnknownMessageError();
@@ -231,6 +238,7 @@ export class MessagePinService extends MessageInteractionBase {
channel_id: channel.id.toString(),
message_id: messageId.toString(),
})
.withThreadScope(channel.isThread())
.withReason(auditLogReason ?? null)
.commit();
}
@@ -6,6 +6,7 @@ import type {IChannelRepositoryAggregate} from '@app/api/channel/repositories/IC
import type {AuthenticatedChannel} from '@app/api/channel/services/AuthenticatedChannel';
import {dispatchChannelEvent} from '@app/api/channel/services/ChannelGatewayDispatch';
import {MessageInteractionBase, type ParsedEmoji} from '@app/api/channel/services/interaction/MessageInteractionBase';
import {assertThreadInteractionAllowed} from '@app/api/channel/services/thread/ThreadInteractionGuards';
import type {IGuildRepositoryAggregate} from '@app/api/guild/repositories/IGuildRepositoryAggregate';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import type {LimitConfigService} from '@app/api/limits/LimitConfigService';
@@ -162,6 +163,7 @@ export class MessageReactionService extends MessageInteractionBase {
const channel = authChannel.channel;
const {guild, hasPermission, checkPermission} = authChannel;
this.ensureTextChannel(channel);
assertThreadInteractionAllowed(authChannel, 'react');
assertGuildMemberCanCommunicate(authChannel.member);
await this.assertMessageHistoryAccess({authChannel, messageId});
if (this.isOperationDisabled(guild, GuildOperations.REACTIONS)) {
@@ -266,6 +268,7 @@ export class MessageReactionService extends MessageInteractionBase {
const channel = authChannel.channel;
const {guild, hasPermission} = authChannel;
this.ensureTextChannel(channel);
assertThreadInteractionAllowed(authChannel, 'react', {ignoreTimeout: true});
await this.assertMessageHistoryAccess({authChannel, messageId});
if (this.isOperationDisabled(guild, GuildOperations.REACTIONS)) {
throw new FeatureTemporarilyDisabledError();
@@ -306,6 +309,7 @@ export class MessageReactionService extends MessageInteractionBase {
const channel = authChannel.channel;
const {guild, hasPermission} = authChannel;
this.ensureTextChannel(channel);
assertThreadInteractionAllowed(authChannel, 'react', {ignoreTimeout: true});
await this.assertMessageHistoryAccess({authChannel, messageId});
if (this.isOperationDisabled(guild, GuildOperations.REACTIONS)) {
throw new FeatureTemporarilyDisabledError();
@@ -338,6 +342,7 @@ export class MessageReactionService extends MessageInteractionBase {
const channel = authChannel.channel;
const {guild, hasPermission} = authChannel;
this.ensureTextChannel(channel);
assertThreadInteractionAllowed(authChannel, 'react', {ignoreTimeout: true});
await this.assertMessageHistoryAccess({authChannel, messageId});
if (this.isOperationDisabled(guild, GuildOperations.REACTIONS)) {
throw new FeatureTemporarilyDisabledError();
@@ -1,7 +1,7 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import fs from 'node:fs';
import {createAttachmentID, type UserID} from '@app/api/BrandedTypes';
import {type ChannelID, createAttachmentID, type UserID} from '@app/api/BrandedTypes';
import {Config} from '@app/api/Config';
import type {AttachmentToProcess} from '@app/api/channel/AttachmentDTOs';
import type {AttachmentUploadTraceRepository} from '@app/api/channel/repositories/message/AttachmentUploadTraceRepository';
@@ -45,6 +45,7 @@ interface ProcessAttachmentParams {
attachment: AttachmentToProcess;
index: number;
uploadUserId: UserID;
uploadChannelId?: ChannelID;
channel?: Channel;
guild?: GuildResponse | null;
member?: GuildMemberResponse | null;
@@ -81,6 +82,7 @@ export class AttachmentProcessingService {
message: Message;
attachments: Array<AttachmentToProcess>;
uploadUserId: UserID;
uploadChannelId?: ChannelID;
channel?: Channel;
guild?: GuildResponse | null;
member?: GuildMemberResponse | null;
@@ -104,6 +106,7 @@ export class AttachmentProcessingService {
attachment,
index,
uploadUserId: params.uploadUserId,
uploadChannelId: params.uploadChannelId,
channel: params.channel,
guild: params.guild,
member: params.member,
@@ -181,7 +184,7 @@ export class AttachmentProcessingService {
const pendingUpload = await this.attachmentUploadTraceRepository.getPendingUpload({
uploadKey: attachment.upload_filename,
userId: params.uploadUserId,
channelId: message.channelId,
channelId: params.uploadChannelId ?? message.channelId,
});
if (!pendingUpload) {
throw InputValidationError.fromCode(
@@ -66,6 +66,7 @@ import {
WebhookTypes,
} from '@fluxer/constants/src/ChannelConstants';
import {GuildFeatures, GuildOperations} from '@fluxer/constants/src/GuildConstants';
import {THREAD_MESSAGE_FLAG_MASK} from '@fluxer/constants/src/ThreadConstants';
import {ContentBlockedError} from '@fluxer/errors/src/domains/content/ContentBlockedError';
import {UnknownGuildError} from '@fluxer/errors/src/domains/guild/UnknownGuildError';
import type {GuildResponse} from '@fluxer/schema/src/domains/guild/GuildResponseSchemas';
@@ -644,7 +645,10 @@ export class CrosspostDeliveryService {
attachments: null,
embeds: null,
sticker_items: null,
flags: MessageFlags.IS_CROSSPOST | MessageFlags.SOURCE_MESSAGE_DELETED,
flags:
MessageFlags.IS_CROSSPOST |
MessageFlags.SOURCE_MESSAGE_DELETED |
(fresh.flags & THREAD_MESSAGE_FLAG_MASK),
edited_timestamp: new Date(),
},
fresh.toRow(),
@@ -8,6 +8,7 @@ import type {IUserRepository} from '@app/api/user/IUserRepository';
import * as EmojiUtils from '@app/api/utils/EmojiUtils';
import {ChannelTypes, GUILD_TEXT_BASED_CHANNEL_TYPES} from '@fluxer/constants/src/ChannelConstants';
import {GuildExplicitContentFilterTypes, GuildFeatures, GuildNSFWLevel} from '@fluxer/constants/src/GuildConstants';
import {THREAD_ONLY_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {SensitiveMediaFilterLevel} from '@fluxer/constants/src/UserConstants';
import type {GuildMemberResponse} from '@fluxer/schema/src/domains/guild/GuildMemberSchemas';
import type {GuildResponse} from '@fluxer/schema/src/domains/guild/GuildResponseSchemas';
@@ -50,7 +51,11 @@ export class MessageContentService {
if (isBot) {
return true;
}
if (channel && GUILD_TEXT_BASED_CHANNEL_TYPES.has(channel.type) && channel.isNsfw) {
if (
channel &&
(GUILD_TEXT_BASED_CHANNEL_TYPES.has(channel.type) || THREAD_ONLY_CHANNEL_TYPES.has(channel.type)) &&
channel.isNsfw
) {
return true;
}
if (channel?.type === ChannelTypes.DM_PERSONAL_NOTES) {
@@ -15,6 +15,7 @@ import type {MessageDispatchService} from '@app/api/channel/services/message/Mes
import {isCrosspostCopy, isOperationDisabled} from '@app/api/channel/services/message/MessageHelpers';
import {assertMessageWithinHistoryCutoff} from '@app/api/channel/services/message/MessageHistoryCutoff';
import type {MessageWriteLock} from '@app/api/channel/services/message/MessageWriteLock';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import {contentModerationService} from '@app/api/infrastructure/ContentModerationService';
import {Logger} from '@app/api/Logger';
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
@@ -58,16 +59,18 @@ export class MessageCrosspostService {
async crosspostMessage({
userId,
viewer,
channelId,
messageId,
requestCache,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
requestCache: RequestCache;
}): Promise<CrosspostMessageResult> {
const authChannel = await this.deps.channelAuthService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.deps.channelAuthService.getChannelAuthenticated({userId, channelId, viewer});
const {channel, guild, member, hasPermission, checkPermission} = authChannel;
if (channel.type !== ChannelTypes.GUILD_ANNOUNCEMENT) {
throw new AnnouncementChannelRequiredError();
@@ -6,9 +6,14 @@ import type {IChannelRepositoryAggregate} from '@app/api/channel/repositories/IC
import type {CrosspostPropagation} from '@app/api/channel/services/message/CrosspostPropagation';
import type {MessageChannelAuthService} from '@app/api/channel/services/message/MessageChannelAuthService';
import type {MessageDispatchService} from '@app/api/channel/services/message/MessageDispatchService';
import {isOperationDisabled, purgeMessageAttachments} from '@app/api/channel/services/message/MessageHelpers';
import {
decrementThreadMessageCount,
isOperationDisabled,
purgeMessageAttachments,
} from '@app/api/channel/services/message/MessageHelpers';
import type {MessageSearchService} from '@app/api/channel/services/message/MessageSearchService';
import type {MessageValidationService} from '@app/api/channel/services/message/MessageValidationService';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import type {GuildAuditLogService} from '@app/api/guild/GuildAuditLogService';
import type {IPurgeQueue} from '@app/api/infrastructure/CachePurgeQueue';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
@@ -20,6 +25,7 @@ import type {Webhook} from '@app/api/models/Webhook';
import {AuditLogActionType} from '@fluxer/constants/src/AuditLogActionType';
import {ChannelTypes, Permissions} from '@fluxer/constants/src/ChannelConstants';
import {GuildOperations} from '@fluxer/constants/src/GuildConstants';
import {TEXT_THREAD_PARENT_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {ValidationErrorCodes} from '@fluxer/constants/src/ValidationErrorCodes';
import {InvalidChannelTypeError} from '@fluxer/errors/src/domains/channel/InvalidChannelTypeError';
import {UnknownMessageError} from '@fluxer/errors/src/domains/channel/UnknownMessageError';
@@ -52,27 +58,39 @@ export class MessageDeleteService {
async deleteMessage({
userId,
viewer,
channelId,
messageId,
skipGuildAuditLog,
auditLogReason,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
requestCache: RequestCache;
skipGuildAuditLog?: boolean;
auditLogReason?: string | null;
}): Promise<void> {
const {channel, guild, hasPermission} = await this.deps.channelAuthService.getChannelAuthenticated({
const {channel, guild, hasPermission, thread} = await this.deps.channelAuthService.getChannelAuthenticated({
userId,
channelId,
viewer,
});
if (isOperationDisabled(guild, GuildOperations.SEND_MESSAGE)) {
throw new FeatureTemporarilyDisabledError();
}
const message = await this.deps.channelRepository.messages.getMessage(channelId, messageId);
if (!message) throw new UnknownMessageError();
if (!message) {
if (
thread?.state.hasStarter &&
TEXT_THREAD_PARENT_CHANNEL_TYPES.has(thread.parent.type) &&
messageId.toString() === channel.id.toString()
) {
throw new MissingPermissionsError();
}
throw new UnknownMessageError();
}
const canDelete = await this.deps.validationService.canDeleteMessage({message, userId, guild, hasPermission});
if (!canDelete) throw new MissingPermissionsError();
if (message.pinnedTimestamp) {
@@ -85,6 +103,7 @@ export class MessageDeleteService {
message.authorId || createUserID(0n),
message.pinnedTimestamp || undefined,
);
await this.decrementThreadMessageCount(channel, [messageId]);
await this.deps.dispatchService.dispatchMessageDelete({channel, messageId, message});
await this.deps.crosspostPropagation.enqueueCrosspostSourceRemoval({
messages: [message],
@@ -106,6 +125,7 @@ export class MessageDeleteService {
.createBuilder(channel.guildId, userId)
.withAction(AuditLogActionType.MESSAGE_DELETE, message.id.toString())
.withMetadata({channel_id: channel.id.toString()})
.withThreadScope(channel.isThread())
.withReason(auditLogReason ?? null)
.commit();
}
@@ -114,14 +134,16 @@ export class MessageDeleteService {
async deleteWebhookMessage({
webhook,
thread,
messageId,
}: {
webhook: Webhook;
thread?: Channel | null;
messageId: MessageID;
requestCache: RequestCache;
}): Promise<void> {
const channelId = webhook.channelId!;
const channel = await this.deps.channelRepository.channelData.findUnique(channelId);
const channelId = thread?.id ?? webhook.channelId!;
const channel = thread ?? (await this.deps.channelRepository.channelData.findUnique(channelId));
if (!channel?.guildId) {
throw new CannotExecuteOnDmError();
}
@@ -140,6 +162,7 @@ export class MessageDeleteService {
message.authorId || createUserID(0n),
message.pinnedTimestamp || undefined,
);
await this.decrementThreadMessageCount(channel, [messageId]);
await this.deps.dispatchService.dispatchMessageDelete({channel, messageId, message});
await this.deps.crosspostPropagation.enqueueCrosspostSourceRemoval({
messages: [message],
@@ -159,13 +182,19 @@ export class MessageDeleteService {
await this.deps.searchService.deleteMessageIndex(messageId);
}
private async decrementThreadMessageCount(channel: Channel, messageIds: Array<MessageID>): Promise<void> {
await decrementThreadMessageCount(this.deps.channelRepository, channel, messageIds);
}
async bulkDeleteMessages({
userId,
viewer,
channelId,
messageIds,
auditLogReason,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageIds: Array<MessageID>;
auditLogReason?: string | null;
@@ -179,6 +208,7 @@ export class MessageDeleteService {
const {channel, guild, checkPermission} = await this.deps.channelAuthService.getChannelAuthenticated({
userId,
channelId,
viewer,
});
if (!guild) throw new CannotExecuteOnDmError();
await checkPermission(Permissions.MANAGE_MESSAGES);
@@ -193,7 +223,12 @@ export class MessageDeleteService {
),
);
await this.deps.channelRepository.messages.bulkDeleteMessages(channelId, messageIds);
await this.deps.dispatchService.dispatchMessageDeleteBulk({channel, messageIds});
const existingIds = existingMessages.map((message) => message.id);
await this.decrementThreadMessageCount(channel, existingIds);
await this.deps.dispatchService.dispatchMessageDeleteBulk({
channel,
messageIds: channel.isThread() ? existingIds : messageIds,
});
await this.deps.crosspostPropagation.enqueueCrosspostSourceRemoval({
messages: existingMessages,
mode: 'source_deleted',
@@ -207,16 +242,25 @@ export class MessageDeleteService {
channel_id: channel.id.toString(),
count: existingMessages.length.toString(),
})
.withThreadScope(channel.isThread())
.withReason(auditLogReason ?? null)
.commit();
}
await this.deps.searchService.deleteMessagesIndex(messageIds);
}
async purgePersonalNotesMessages({userId, channelId}: {userId: UserID; channelId: ChannelID}): Promise<{
async purgePersonalNotesMessages({
userId,
channelId,
viewer,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
}): Promise<{
deletedCount: number;
}> {
const {channel} = await this.deps.channelAuthService.getChannelAuthenticated({userId, channelId});
const {channel} = await this.deps.channelAuthService.getChannelAuthenticated({userId, channelId, viewer});
if (
channel.type !== ChannelTypes.DM_PERSONAL_NOTES ||
!this.deps.channelAuthService.isPersonalNotesChannel({userId, channelId})
@@ -252,8 +296,16 @@ export class MessageDeleteService {
guildId: GuildID;
seconds: number;
}): Promise<void> {
const channels = await this.deps.channelRepository.channelData.listGuildChannels(guildId);
const cutoffTimestamp = Date.now() - seconds * ms('1 second');
const guildChannels = await this.deps.channelRepository.channelData.listGuildChannels(guildId, 'complete');
const threadIds = await this.deps.channelRepository.threads.listGuildThreadIds(guildId, {
parents: guildChannels,
activeSince: new Date(cutoffTimestamp),
});
const channels =
threadIds.length > 0
? [...guildChannels, ...(await this.deps.channelRepository.channelData.listChannels(threadIds))]
: guildChannels;
const cutoffSnowflake = createMessageID(createSnowflakeFromTimestamp(cutoffTimestamp));
await Promise.all(
channels.map(async (channel: Channel) => {
@@ -276,6 +328,7 @@ export class MessageDeleteService {
),
);
await this.deps.channelRepository.messages.bulkDeleteMessages(channel.id, messageIds);
await this.decrementThreadMessageCount(channel, messageIds);
await this.deps.dispatchService.dispatchMessageDeleteBulk({channel, messageIds});
await this.deps.crosspostPropagation.enqueueCrosspostSourceRemoval({
messages: userMessages,
@@ -16,6 +16,8 @@ import type {MessageProcessingService} from '@app/api/channel/services/message/M
import type {MessageSearchService} from '@app/api/channel/services/message/MessageSearchService';
import type {MessageValidationService} from '@app/api/channel/services/message/MessageValidationService';
import type {MessageWriteLock} from '@app/api/channel/services/message/MessageWriteLock';
import {assertThreadInteractionAllowed} from '@app/api/channel/services/thread/ThreadInteractionGuards';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import {Logger} from '@app/api/Logger';
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
import type {Message} from '@app/api/models/Message';
@@ -57,12 +59,14 @@ export class MessageEditService {
async editMessage({
userId,
viewer,
channelId,
messageId,
data,
requestCache,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
data: MessageUpdateRequest;
@@ -71,6 +75,7 @@ export class MessageEditService {
const authChannel = await this.deps.channelAuthService.getChannelAuthenticated({
userId,
channelId,
viewer,
});
const {channel, guild, hasPermission, member} = authChannel;
const hasNewAttachments =
@@ -99,6 +104,7 @@ export class MessageEditService {
if (message.authorId === userId) {
assertGuildMemberCanCommunicate(member);
}
assertThreadInteractionAllowed(authChannel, 'edit', {ignoreTimeout: message.authorId !== userId});
if (data.message_snapshots !== undefined) {
throw new MissingPermissionsError();
}
@@ -1,8 +1,9 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {AttachmentID, ChannelID, UserID} from '@app/api/BrandedTypes';
import type {AttachmentID, ChannelID, MessageID, UserID} from '@app/api/BrandedTypes';
import {createAttachmentID, userIdToChannelId} from '@app/api/BrandedTypes';
import {Config} from '@app/api/Config';
import type {IChannelRepositoryAggregate} from '@app/api/channel/repositories/IChannelRepositoryAggregate';
import {forEachEmbedMedia} from '@app/api/channel/services/message/CrosspostEmbedObjects';
import type {
MessageSnapshot as CassandraMessageSnapshot,
@@ -18,6 +19,7 @@ import type {LimitConfigService} from '@app/api/limits/LimitConfigService';
import {resolveLimitSafe} from '@app/api/limits/LimitConfigUtils';
import {createLimitMatchContext} from '@app/api/limits/LimitMatchContextBuilder';
import {Attachment} from '@app/api/models/Attachment';
import type {Channel} from '@app/api/models/Channel';
import type {Message} from '@app/api/models/Message';
import {MessageSnapshot as MessageSnapshotModel} from '@app/api/models/MessageSnapshot';
import type {User} from '@app/api/models/User';
@@ -469,3 +471,19 @@ export function isMessageEmpty(message: Message, excludingAttachments = false):
export function collectMessageAttachments(message: Message): Array<Attachment> {
return [...message.attachments, ...message.messageSnapshots.flatMap((snapshot) => snapshot.attachments)];
}
export function countedThreadMessages(channel: Channel, messageIds: Array<MessageID>): number {
if (!channel.isThread()) return 0;
const starterId = channel.id.toString();
return messageIds.filter((messageId) => messageId.toString() !== starterId).length;
}
export async function decrementThreadMessageCount(
channelRepository: IChannelRepositoryAggregate,
channel: Channel | null,
messageIds: Array<MessageID>,
): Promise<void> {
if (!channel) return;
const count = countedThreadMessages(channel, messageIds);
if (count > 0) await channelRepository.threads.adjustMessageCount(channel.id, -count);
}
@@ -19,6 +19,7 @@ import {
keepOwnedEmbedAttachments,
} from '@app/api/channel/services/message/MessageHelpers';
import {MessageStickerService} from '@app/api/channel/services/message/MessageStickerService';
import {resolveNsfwScopeChannel} from '@app/api/channel/utils/ThreadNsfwScope';
import {getContentMessage} from '@app/api/content_i18n/ContentI18n';
import type {
MessageAttachment,
@@ -93,6 +94,7 @@ interface CreateMessageParams {
embeds?: Array<RichEmbedRequest>;
attachments?: Array<AttachmentToProcess>;
attachmentUploadUserId?: UserID;
uploadChannelId?: ChannelID;
processedAttachments?: Array<MessageAttachment>;
stickerIds?: Array<StickerID>;
messageReference?: MessageReference;
@@ -116,6 +118,7 @@ interface CreateMessageParams {
processedEmbeds?: Array<MessageEmbed>;
processedStickerItems?: Array<MessageStickerItem>;
skipDeferredEmbeds?: boolean;
threadInsert?: boolean;
}
export class MessagePersistenceService {
@@ -151,6 +154,10 @@ export class MessagePersistenceService {
this.attachmentDecayService = new AttachmentDecayService();
}
private nsfwScopeChannel(channel: Channel): Promise<Channel> {
return resolveNsfwScopeChannel(channel, (channelId) => this.channelRepository.channelData.findUnique(channelId));
}
getEmbedAttachmentResolver(): MessageEmbedAttachmentResolver {
return this.embedAttachmentResolver;
}
@@ -169,7 +176,7 @@ export class MessagePersistenceService {
const isBot = params.user?.isBot ?? false;
const isBugHunterBot = isBot && ((params.user?.flags ?? 0n) & UserFlags.BUG_HUNTER) !== 0n;
const isNSFWAllowed = this.contentService.isNSFWContentAllowed({
channel: params.channel,
channel: params.channel ? await this.nsfwScopeChannel(params.channel) : undefined,
guild: params.guild,
member: params.member,
isBot,
@@ -248,7 +255,11 @@ export class MessagePersistenceService {
has_reaction: false,
version: 1,
};
const message = await this.channelRepository.messages.upsertMessage(messageRowData, null);
const message = await this.channelRepository.messages.upsertMessage(
messageRowData,
null,
params.threadInsert ? {isInsert: true} : undefined,
);
const enqueueDeferredEmbeds = await this.runPostPersistenceOperations({
message,
params,
@@ -297,6 +308,7 @@ export class MessagePersistenceService {
} as Message,
attachments: params.attachments,
uploadUserId,
uploadChannelId: params.uploadChannelId,
channel: params.channel,
guild: params.guild,
member: params.member,
@@ -367,6 +379,7 @@ export class MessagePersistenceService {
mentionCount: 0,
implicit: {unreadThrough: params.user ? (params.channel?.lastMessageId ?? null) : null},
emitGateway: false,
...(params.threadInsert ? {capable: true, channel: params.channel ?? null} : {}),
}),
);
}
@@ -395,7 +408,7 @@ export class MessagePersistenceService {
throw InputValidationError.fromCode('message', ValidationErrorCodes.MESSAGES_WITH_SNAPSHOTS_CANNOT_BE_EDITED);
}
const isNSFWAllowed = this.contentService.isNSFWContentAllowed({
channel,
channel: await this.nsfwScopeChannel(channel),
guild,
member,
isBot: params.isBot,
@@ -715,6 +728,7 @@ export class MessagePersistenceService {
guildId: params.guildId ?? null,
allowEmbeds: false,
messageReference: params.messageReference,
threadInsert: true,
mentionData: {
flags: 0,
mentionUserIds: params.mentionUserIds ?? [],
@@ -10,6 +10,7 @@ import {
import type {IChannelRepository} from '@app/api/channel/IChannelRepository';
import type {MessageRequest, MessageUpdateRequest} from '@app/api/channel/MessageTypes';
import {normalizeMessageRequestPayload} from '@app/api/channel/services/message/MessageRequestCompatibility';
import {SYSTEM_THREAD_VIEWER, viewerFromCtx} from '@app/api/experiment/ChannelThreadsGate';
import type {GuildService} from '@app/api/guild/services/GuildService';
import type {LimitConfigService} from '@app/api/limits/LimitConfigService';
import {resolveLimitSafe} from '@app/api/limits/LimitConfigUtils';
@@ -159,6 +160,7 @@ export async function parseMultipartMessageData(
.get('channelService')
.attachments.uploadFormDataAttachments({
userId: user.id,
viewer: options?.actor === 'webhook' ? SYSTEM_THREAD_VIEWER : viewerFromCtx(ctx),
channelId,
clientIp,
files: filesWithIndices,
@@ -0,0 +1,163 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {createChannelID, createGuildID, createMessageID, createUserID} from '@app/api/BrandedTypes';
import type {ChannelService} from '@app/api/channel/services/ChannelService';
import type {CrosspostSourceService} from '@app/api/channel/services/message/CrosspostSourceService';
import {MessageRequestService} from '@app/api/channel/services/message/MessageRequestService';
import type {MessageResponseDataService} from '@app/api/channel/services/message/MessageResponseDataService';
import {setCassandraQueryExecutorForTesting, upsertOne} from '@app/api/database/CassandraQueryExecution';
import type {CassandraParams, PreparedQuery} from '@app/api/database/CassandraTypes';
import {
clearChannelThreadsTaintCacheForTesting,
syncChannelThreadsConfig,
type ThreadViewer,
} from '@app/api/experiment/ChannelThreadsGate';
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
import {GuildThreadState} from '@app/api/Tables';
import {InMemoryCassandraQueryExecutor} from '@app/api/test/InMemoryCassandraQueryExecutor';
import {ChannelTypes, MessageTypes} from '@fluxer/constants/src/ChannelConstants';
import {ServerMessageFlags} from '@fluxer/constants/src/ThreadConstants';
import {ChannelThreadsConfigSchema} from '@fluxer/schema/src/domains/admin/ChannelThreadsSchemas';
import type {MessageResponse} from '@fluxer/schema/src/domains/message/MessageResponseSchemas';
import {afterEach, beforeEach, describe, expect, it, vi} from 'vitest';
class MarkerCountingExecutor extends InMemoryCassandraQueryExecutor {
markerReads = 0;
override async executeQuery<T = Record<string, unknown>, P extends CassandraParams = CassandraParams>(
query: PreparedQuery<P>,
): Promise<Array<T>> {
if (query.cql.trimStart().toUpperCase().startsWith('SELECT') && query.cql.includes('guild_thread_state')) {
this.markerReads++;
}
return super.executeQuery<T>(query);
}
}
const GUILD = createGuildID(100n);
const CHANNEL = createChannelID(200n);
const USER = createUserID(10n);
const VIEWER: ThreadViewer = {kind: 'user', userId: USER, bot: false, capable: false};
function message(overrides: Partial<MessageResponse> = {}): MessageResponse {
return {
id: '300',
channel_id: CHANNEL.toString(),
type: MessageTypes.DEFAULT,
flags: 0,
content: 'hello',
mentions: [],
mention_roles: [],
...overrides,
} as MessageResponse;
}
function build(responses: Array<MessageResponse>) {
const listMessages = vi.fn(async (_params: {threadsMask?: boolean}) => responses);
const getMessage = vi.fn(async (_params: {threadsMask?: boolean}) => responses[0] ?? null);
const channelService = {
messages: {
retrieval: {
getResponseAccess: async () => ({
access: {sourceGuildId: GUILD, messageHistoryCutoff: null, canReadMessageHistory: true},
authChannel: {channel: {guildId: GUILD, isThreadOnly: () => false}},
}),
threadResponses: {
shape: async (params: {responses: Array<MessageResponse>}) => params.responses,
getStarter: async () => null,
},
},
},
} as unknown as ChannelService;
const service = new MessageRequestService(
channelService,
{listMessages, getMessage} as unknown as MessageResponseDataService,
{} as CrosspostSourceService,
);
return {service, listMessages, getMessage};
}
function list(service: MessageRequestService) {
return service.listMessages({
userId: USER,
viewer: VIEWER,
channelId: CHANNEL,
query: {limit: 50},
requestCache: {} as RequestCache,
});
}
function get(service: MessageRequestService) {
return service.getMessage({
userId: USER,
viewer: VIEWER,
channelId: CHANNEL,
messageId: createMessageID(300n),
requestCache: {} as RequestCache,
});
}
async function taint(): Promise<void> {
await upsertOne(
GuildThreadState.upsertAll({
guild_id: GUILD,
first_active_at: new Date(),
perms_seeded_at: null,
search_backfilled_at: null,
}),
);
}
describe('MessageRequestService thread masking', () => {
let executor: MarkerCountingExecutor;
beforeEach(() => {
executor = new MarkerCountingExecutor();
setCassandraQueryExecutorForTesting(executor);
clearChannelThreadsTaintCacheForTesting();
const raw = JSON.stringify(ChannelThreadsConfigSchema.parse({enabled: false, ever_enabled: true}));
syncChannelThreadsConfig(raw, (value) => ChannelThreadsConfigSchema.parse(JSON.parse(value ?? '{}')));
});
afterEach(() => {
syncChannelThreadsConfig(null, () => ChannelThreadsConfigSchema.parse({}));
setCassandraQueryExecutorForTesting(null);
});
it('never reads the marker for a page without thread data', async () => {
const {service, listMessages, getMessage} = build([message()]);
await list(service);
await get(service);
expect(executor.markerReads).toBe(0);
expect(listMessages).toHaveBeenCalledTimes(1);
expect(listMessages.mock.calls[0]![0].threadsMask).toBeUndefined();
expect(getMessage).toHaveBeenCalledTimes(1);
expect(getMessage.mock.calls[0]![0].threadsMask).toBeUndefined();
});
it('refetches masked when a tainted guild page carries thread data', async () => {
await taint();
const {service, listMessages, getMessage} = build([message({flags: ServerMessageFlags.HAS_THREAD})]);
await list(service);
await get(service);
expect(executor.markerReads).toBe(1);
expect(listMessages.mock.calls.map(([params]) => params.threadsMask)).toEqual([undefined, true]);
expect(getMessage.mock.calls.map(([params]) => params.threadsMask)).toEqual([undefined, true]);
});
it('treats a mentioned forum as thread data', async () => {
await taint();
const {service, listMessages} = build([
message({mention_channels: [{id: '400', name: 'forum', type: ChannelTypes.GUILD_FORUM}]}),
]);
await list(service);
expect(listMessages.mock.calls.map(([params]) => params.threadsMask)).toEqual([undefined, true]);
});
it('keeps raw data in an untainted guild', async () => {
const {service, listMessages} = build([message({flags: ServerMessageFlags.HAS_THREAD})]);
await list(service);
expect(executor.markerReads).toBe(1);
expect(listMessages).toHaveBeenCalledTimes(1);
});
});
@@ -1,22 +1,39 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {ChannelID, MessageID, UserID} from '@app/api/BrandedTypes';
import type {ChannelID, GuildID, MessageID, UserID} from '@app/api/BrandedTypes';
import type {MessageRequest, MessageUpdateRequest} from '@app/api/channel/MessageTypes';
import type {AuthenticatedChannel} from '@app/api/channel/services/AuthenticatedChannel';
import type {ChannelService} from '@app/api/channel/services/ChannelService';
import type {CrosspostSourceService} from '@app/api/channel/services/message/CrosspostSourceService';
import {isPersonalNotesChannel} from '@app/api/channel/services/message/MessageHelpers';
import type {MessageResponseDataService} from '@app/api/channel/services/message/MessageResponseDataService';
import {carriesThreadArtifact, maskThreadArtifactsFor} from '@app/api/channel/services/message/ThreadMessageResponses';
import {everEnabled, isTainted, type ThreadViewer, viewerActive} from '@app/api/experiment/ChannelThreadsGate';
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
import type {User} from '@app/api/models/User';
import {mapWithConcurrency} from '@app/api/utils/ConcurrencyUtils';
import {THREAD_FEATURE_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {InvalidChannelTypeError} from '@fluxer/errors/src/domains/channel/InvalidChannelTypeError';
import {UnclaimedAccountCannotSendMessagesError} from '@fluxer/errors/src/domains/channel/UnclaimedAccountCannotSendMessagesError';
import {UnknownMessageError} from '@fluxer/errors/src/domains/channel/UnknownMessageError';
import type {CrosspostSourceResponse} from '@fluxer/schema/src/domains/message/CrosspostSourceSchemas';
import type {
BulkMessageFetchResponse,
MessageChannelMentionResponse,
MessageResponse,
} from '@fluxer/schema/src/domains/message/MessageResponseSchemas';
function mentionsThreadChannel(mentions: ReadonlyArray<MessageChannelMentionResponse> | null | undefined): boolean {
return mentions?.some((mention) => THREAD_FEATURE_CHANNEL_TYPES.has(mention.type)) ?? false;
}
function carriesMaskableThreadData(response: MessageResponse): boolean {
if (carriesThreadArtifact(response) || mentionsThreadChannel(response.mention_channels)) return true;
if (response.message_snapshots?.some((snapshot) => mentionsThreadChannel(snapshot.mention_channels))) return true;
const referenced = response.referenced_message;
return referenced != null && carriesMaskableThreadData(referenced);
}
export class MessageRequestService {
constructor(
private readonly channelService: ChannelService,
@@ -26,6 +43,7 @@ export class MessageRequestService {
async listMessages(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
query: {
limit: number;
@@ -35,11 +53,14 @@ export class MessageRequestService {
};
requestCache: RequestCache;
}): Promise<Array<MessageResponse>> {
const access = await this.channelService.messages.retrieval.getResponseAccessContext({
const retrieval = this.channelService.messages.retrieval;
const {access, authChannel} = await retrieval.getResponseAccess({
viewer: params.viewer,
userId: params.userId,
channelId: params.channelId,
});
return this.responseDataService.listMessages({
if (authChannel.channel.isThreadOnly()) throw new InvalidChannelTypeError();
const request = {
userId: params.userId,
channelId: params.channelId,
limit: params.query.limit,
@@ -47,11 +68,34 @@ export class MessageRequestService {
after: params.query.after,
around: params.query.around,
access,
};
let responses = await this.responseDataService.listMessages(request);
if (await this.needsThreadsMask(params.viewer, authChannel.channel.guildId, responses)) {
responses = await this.responseDataService.listMessages({...request, threadsMask: true});
}
return retrieval.threadResponses.shape({
viewer: params.viewer,
userId: params.userId,
authChannel,
responses,
requestCache: params.requestCache,
query: params.query,
});
}
private async needsThreadsMask(
viewer: ThreadViewer,
guildId: GuildID | null,
responses: Array<MessageResponse>,
): Promise<boolean> {
if (guildId === null || !everEnabled() || viewerActive(viewer, guildId)) return false;
if (!responses.some(carriesMaskableThreadData)) return false;
return isTainted(guildId);
}
async listMessagesBulk(params: {
userId: UserID;
viewer: ThreadViewer;
requests: Array<{
channelId: ChannelID;
query: {
@@ -67,6 +111,7 @@ export class MessageRequestService {
channel_id: request.channelId.toString(),
messages: await this.listMessages({
userId: params.userId,
viewer: params.viewer,
channelId: request.channelId,
query: request.query,
requestCache: params.requestCache,
@@ -77,29 +122,55 @@ export class MessageRequestService {
async getMessage(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
requestCache: RequestCache;
}): Promise<MessageResponse> {
const access = await this.channelService.messages.retrieval.getResponseAccessContext({
const retrieval = this.channelService.messages.retrieval;
const {access, authChannel} = await retrieval.getResponseAccess({
viewer: params.viewer,
userId: params.userId,
channelId: params.channelId,
messageId: params.messageId,
});
const response = await this.responseDataService.getMessage({
const starter = await retrieval.threadResponses.getStarter({
viewer: params.viewer,
userId: params.userId,
authChannel,
messageId: params.messageId,
requestCache: params.requestCache,
});
if (starter) return starter;
const request = {
userId: params.userId,
channelId: params.channelId,
messageId: params.messageId,
access,
});
};
let response = await this.responseDataService.getMessage(request);
if (response !== null && (await this.needsThreadsMask(params.viewer, authChannel.channel.guildId, [response]))) {
response = await this.responseDataService.getMessage({...request, threadsMask: true});
}
if (response === null) {
throw new UnknownMessageError();
}
return response;
const [shaped] = await retrieval.threadResponses.shape({
viewer: params.viewer,
userId: params.userId,
authChannel,
responses: [response],
requestCache: params.requestCache,
});
if (!shaped) {
throw new UnknownMessageError();
}
return shaped;
}
async getCrosspostSource(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
requestCache: RequestCache;
@@ -110,6 +181,7 @@ export class MessageRequestService {
async sendMessage(params: {
user: User;
viewer: ThreadViewer;
channelId: ChannelID;
data: MessageRequest;
requestCache: RequestCache;
@@ -121,53 +193,109 @@ export class MessageRequestService {
throw new UnclaimedAccountCannotSendMessagesError();
}
const {message, authChannel} = await this.channelService.messages.send.sendMessage({
viewer: params.viewer,
user: params.user,
channelId: params.channelId,
data: params.data,
requestCache: params.requestCache,
});
const access = await this.channelService.messages.retrieval.getResponseAccessContext({
viewer: params.viewer,
userId: params.user.id,
channelId: params.channelId,
authChannel,
});
return this.responseDataService.buildMessage({
const response = await this.responseDataService.buildMessage({
userId: params.user.id,
message,
access: {...access, messageHistoryCutoff: null, canReadMessageHistory: true},
nonce: params.data.nonce,
tts: params.data.tts ?? false,
});
const [shaped] = maskThreadArtifactsFor(params.viewer, authChannel.channel.guildId, [response]);
return shaped ?? response;
}
async crosspostMessage(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
requestCache: RequestCache;
}): Promise<MessageResponse> {
const {message, authChannel} = await this.channelService.messages.crosspost.crosspostMessage(params);
const access = await this.channelService.messages.retrieval.getResponseAccessContext({
const retrieval = this.channelService.messages.retrieval;
const access = await retrieval.getResponseAccessContext({
viewer: params.viewer,
userId: params.userId,
channelId: params.channelId,
messageId: message.id,
authChannel,
});
return this.responseDataService.buildMessage({
const response = await this.responseDataService.buildMessage({
userId: params.userId,
message,
access,
});
const [shaped] = await retrieval.threadResponses.shape({
viewer: params.viewer,
userId: params.userId,
authChannel,
responses: [response],
requestCache: params.requestCache,
});
return shaped ?? response;
}
async validateForumStarter(params: {
user: User;
parentAuth: AuthenticatedChannel;
data: MessageRequest;
}): Promise<void> {
if (params.user.isUnclaimedAccount()) throw new UnclaimedAccountCannotSendMessagesError();
await this.channelService.messages.send.validateForumStarter(params);
}
async sendForumStarter(params: {
user: User;
viewer: ThreadViewer;
parentAuth: AuthenticatedChannel;
threadId: ChannelID;
data: MessageRequest;
requestCache: RequestCache;
}): Promise<MessageResponse> {
if (params.user.isUnclaimedAccount()) throw new UnclaimedAccountCannotSendMessagesError();
const {message, authChannel} = await this.channelService.messages.send.sendMessage({
viewer: params.viewer,
user: params.user,
channelId: params.threadId,
data: params.data,
requestCache: params.requestCache,
forumStarter: {parentAuth: params.parentAuth},
});
const access = await this.channelService.messages.retrieval.getResponseAccessContext({
viewer: params.viewer,
userId: params.user.id,
channelId: params.threadId,
authChannel,
});
return this.responseDataService.buildMessage({
userId: params.user.id,
message,
access: {...access, messageHistoryCutoff: null, canReadMessageHistory: true},
});
}
async editMessage(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
data: MessageUpdateRequest;
requestCache: RequestCache;
}): Promise<MessageResponse> {
const {message, authChannel} = await this.channelService.messages.edit.editMessage({
viewer: params.viewer,
userId: params.userId,
channelId: params.channelId,
messageId: params.messageId,
@@ -175,15 +303,18 @@ export class MessageRequestService {
requestCache: params.requestCache,
});
const access = await this.channelService.messages.retrieval.getResponseAccessContext({
viewer: params.viewer,
userId: params.userId,
channelId: params.channelId,
messageId: message.id,
authChannel,
});
return this.responseDataService.buildMessage({
const response = await this.responseDataService.buildMessage({
userId: params.userId,
message,
access,
});
const [shaped] = maskThreadArtifactsFor(params.viewer, authChannel.channel.guildId, [response]);
return shaped ?? response;
}
}
@@ -8,6 +8,7 @@ import {Logger} from '@app/api/Logger';
import type {Channel} from '@app/api/models/Channel';
import type {Message} from '@app/api/models/Message';
import {isJsonRecord, parseJsonRecord, parseJsonWithGuard} from '@app/api/utils/JsonBoundaryUtils';
import {MessageTypes} from '@fluxer/constants/src/ChannelConstants';
import type {MessageResponse} from '@fluxer/schema/src/domains/message/MessageResponseSchemas';
import type {INatsConnectionManager} from '@pkgs/nats/src/INatsConnectionManager';
import {NatsConnectionManager} from '@pkgs/nats/src/NatsConnectionManager';
@@ -92,8 +93,10 @@ export class MessageResponseDataService {
after?: MessageID;
around?: MessageID;
access: MessageResponseAccessContext;
threadsMask?: boolean;
}): Promise<Array<MessageResponse>> {
const response = await this.request({
...(params.threadsMask ? {threads_mask: true, exclude_types: [MessageTypes.THREAD_CREATED]} : {}),
op: 'ListResponses',
channel_id: params.channelId.toString(),
viewer_user_id: params.userId.toString(),
@@ -137,8 +140,10 @@ export class MessageResponseDataService {
access: MessageResponseAccessContext;
nonce?: string;
tts?: boolean;
threadsMask?: boolean;
}): Promise<MessageResponse | null> {
const response = await this.request({
...(params.threadsMask ? {threads_mask: true} : {}),
op: 'GetResponseById',
channel_id: params.channelId.toString(),
message_id: params.messageId.toString(),
@@ -9,6 +9,7 @@ import type {MessageChannelAuthService} from '@app/api/channel/services/message/
import type {MessageProcessingService} from '@app/api/channel/services/message/MessageProcessingService';
import {MessageRetrievalService} from '@app/api/channel/services/message/MessageRetrievalService';
import type {MessageSearchService} from '@app/api/channel/services/message/MessageSearchService';
import {SYSTEM_THREAD_VIEWER} from '@app/api/experiment/ChannelThreadsGate';
import type {UserCacheService} from '@app/api/infrastructure/UserCacheService';
import {Message} from '@app/api/models/Message';
import type {IUserRepository} from '@app/api/user/IUserRepository';
@@ -121,6 +122,7 @@ describe('MessageRetrievalService.getMessagesByIds', () => {
});
const result = await service.getMessagesByIds({
viewer: SYSTEM_THREAD_VIEWER,
userId: VIEWER_ID,
channelId: CHANNEL_ID,
messageIds: messages.map((message) => message.id),
@@ -141,6 +143,7 @@ describe('MessageRetrievalService.getMessagesByIds', () => {
});
const result = await service.getMessagesByIds({
viewer: SYSTEM_THREAD_VIEWER,
userId: VIEWER_ID,
channelId: CHANNEL_ID,
messageIds: [beforeCutoff.id, afterCutoff.id],
@@ -160,6 +163,7 @@ describe('MessageRetrievalService.getMessagesByIds', () => {
});
const result = await service.getMessagesByIds({
viewer: SYSTEM_THREAD_VIEWER,
userId: VIEWER_ID,
channelId: CHANNEL_ID,
messageIds: [message.id],
@@ -177,6 +181,7 @@ describe('MessageRetrievalService.getMessagesByIds', () => {
});
const result = await service.getMessagesByIds({
viewer: SYSTEM_THREAD_VIEWER,
userId: VIEWER_ID,
channelId: CHANNEL_ID,
messageIds: [message.id],
@@ -196,6 +201,7 @@ describe('MessageRetrievalService.getMessagesByIds', () => {
});
const result = await service.getMessagesByIds({
viewer: SYSTEM_THREAD_VIEWER,
userId: VIEWER_ID,
channelId: CHANNEL_ID,
messageIds: [present.id, deleted.id],
@@ -18,6 +18,9 @@ import {
type MessageResponseAccessContext,
} from '@app/api/channel/services/message/MessageResponseDataService';
import type {MessageSearchService} from '@app/api/channel/services/message/MessageSearchService';
import {maskThreadArtifactsFor, ThreadMessageResponses} from '@app/api/channel/services/message/ThreadMessageResponses';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import {maskChannelResponseThreadBits} from '@app/api/guild/services/ThreadPermissionBits';
import type {UserCacheService} from '@app/api/infrastructure/UserCacheService';
import type {RequestCache} from '@app/api/middleware/RequestCacheMiddleware';
import type {Channel} from '@app/api/models/Channel';
@@ -65,17 +68,40 @@ export class MessageRetrievalService {
return this.isMessageAfterCutoff(messageId, cutoff);
}
private threadResponsesInstance: ThreadMessageResponses | undefined;
get threadResponses(): ThreadMessageResponses {
this.threadResponsesInstance ??= new ThreadMessageResponses(
this.channelRepository,
createMessageResponseDataService(),
this.userCacheService,
);
return this.threadResponsesInstance;
}
async getResponseAccessContext(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId?: MessageID;
authChannel?: AuthenticatedChannel;
}): Promise<MessageResponseAccessContext> {
return (await this.getResponseAccess(params)).access;
}
async getResponseAccess(params: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId?: MessageID;
authChannel?: AuthenticatedChannel;
}): Promise<{access: MessageResponseAccessContext; authChannel: AuthenticatedChannel}> {
const authChannel =
params.authChannel ??
(await this.channelAuthService.getChannelAuthenticated({
userId: params.userId,
channelId: params.channelId,
viewer: params.viewer,
}));
if (params.messageId && !(await this.canAccessMessage(authChannel, params.messageId))) {
throw new UnknownMessageError();
@@ -83,22 +109,27 @@ export class MessageRetrievalService {
const canReadMessageHistory =
!authChannel.guild || (await authChannel.hasPermission(Permissions.READ_MESSAGE_HISTORY));
return {
sourceGuildId: authChannel.channel.guildId,
messageHistoryCutoff: canReadMessageHistory ? null : (authChannel.guild?.message_history_cutoff ?? null),
canReadMessageHistory,
authChannel,
access: {
sourceGuildId: authChannel.channel.guildId,
messageHistoryCutoff: canReadMessageHistory ? null : (authChannel.guild?.message_history_cutoff ?? null),
canReadMessageHistory,
},
};
}
async getMessage({
userId,
viewer,
channelId,
messageId,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageId: MessageID;
}): Promise<Message> {
const authChannel = await this.channelAuthService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.channelAuthService.getChannelAuthenticated({userId, channelId, viewer});
if (!(await this.canAccessMessage(authChannel, messageId))) {
throw new UnknownMessageError();
}
@@ -113,14 +144,16 @@ export class MessageRetrievalService {
async getMessagesByIds({
userId,
viewer,
channelId,
messageIds,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
messageIds: Array<MessageID>;
}): Promise<Map<string, Message>> {
const authChannel = await this.channelAuthService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.channelAuthService.getChannelAuthenticated({userId, channelId, viewer});
const canReadMessageHistory =
!authChannel.guild || (await authChannel.hasPermission(Permissions.READ_MESSAGE_HISTORY));
const cutoff = authChannel.guild?.message_history_cutoff ?? null;
@@ -141,16 +174,18 @@ export class MessageRetrievalService {
async searchMessages({
userId,
viewer,
channelId,
searchParams,
requestCache,
}: {
userId: UserID;
viewer: ThreadViewer;
channelId: ChannelID;
searchParams: MessageSearchRequest;
requestCache: RequestCache;
}): Promise<MessageSearchResponse> {
const authChannel = await this.channelAuthService.getChannelAuthenticated({userId, channelId});
const authChannel = await this.channelAuthService.getChannelAuthenticated({userId, channelId, viewer});
const {channel} = authChannel;
const hasReadHistory = !authChannel.guild || (await authChannel.hasPermission(Permissions.READ_MESSAGE_HISTORY));
if (!hasReadHistory) {
@@ -208,11 +243,12 @@ export class MessageRetrievalService {
messages: result.messages,
access,
});
const messageResponses = builtMessages.map(
const messageResponses = maskThreadArtifactsFor(viewer, channel.guildId, builtMessages).map(
({referenced_message: _referencedMessage, ...searchMessage}) => searchMessage,
);
return {
channels: messageResponses.length > 0 ? [await this.mapSearchChannelResponse(channel, userId, requestCache)] : [],
channels:
messageResponses.length > 0 ? await this.mapSearchChannelResponse(channel, userId, viewer, requestCache) : [],
messages: messageResponses,
total: hasReadHistory ? result.total : messageResponses.length,
hits_per_page: hitsPerPage,
@@ -220,13 +256,19 @@ export class MessageRetrievalService {
};
}
private async mapSearchChannelResponse(channel: Channel, userId: UserID, requestCache: RequestCache) {
return mapChannelToResponse({
private async mapSearchChannelResponse(
channel: Channel,
userId: UserID,
viewer: ThreadViewer,
requestCache: RequestCache,
) {
const response = await mapChannelToResponse({
channel,
currentUserId: userId,
userCacheService: this.userCacheService,
requestCache,
});
return maskChannelResponseThreadBits(channel.guildId, viewer, [response]);
}
private async channelNeedsIndexing(channel: Channel, channelId: ChannelID): Promise<boolean> {
@@ -8,6 +8,7 @@ import type {IMessageSearchService} from '@app/api/search/IMessageSearchService'
import {deleteMessageSearchDocuments} from '@app/api/search/MessageSearchIndexCleanup';
import type {IUserRepository} from '@app/api/user/IUserRepository';
import type {WorkerTaskName} from '@app/api/worker/WorkerLaneConfig';
import {MessageTypes} from '@fluxer/constants/src/ChannelConstants';
import type {MessageSearchFilters} from '@fluxer/schema/src/contracts/search/SearchDocumentTypes';
import type {MessageSearchRequest} from '@fluxer/schema/src/domains/message/MessageRequestSchemas';
import type {IWorkerService} from '@pkgs/worker/src/contracts/IWorkerService';
@@ -26,6 +27,10 @@ function getMessageIndexServices(options: MessageSearchIndexOptions = {}): Array
return services;
}
export function isMessageSearchIndexable(message: Message): boolean {
return message.type !== MessageTypes.THREAD_CREATED && message.type !== MessageTypes.THREAD_STARTER_MESSAGE;
}
export class MessageSearchService {
constructor(
private userRepository: IUserRepository,
@@ -33,6 +38,7 @@ export class MessageSearchService {
) {}
async indexMessage(message: Message, authorIsBot: boolean, options?: MessageSearchIndexOptions): Promise<void> {
if (!isMessageSearchIndexable(message)) return;
try {
const searchServices = getMessageIndexServices(options);
if (searchServices.length === 0) {
@@ -54,6 +60,7 @@ export class MessageSearchService {
}
async updateMessageIndex(message: Message, options?: MessageSearchIndexOptions): Promise<void> {
if (!isMessageSearchIndexable(message)) return;
try {
const searchServices = getMessageIndexServices(options);
if (searchServices.length === 0) {
@@ -2,6 +2,7 @@
import type {AttachmentID, ChannelID, GuildID, MessageID, RoleID, UserID} from '@app/api/BrandedTypes';
import {
channelIdToMessageId,
createAttachmentID,
createChannelID,
createGuildID,
@@ -35,8 +36,12 @@ import type {MessageProcessingService} from '@app/api/channel/services/message/M
import type {MessageSearchService} from '@app/api/channel/services/message/MessageSearchService';
import type {MessageValidationService} from '@app/api/channel/services/message/MessageValidationService';
import type {MessageWriteLock} from '@app/api/channel/services/message/MessageWriteLock';
import type {ThreadMessageActivity} from '@app/api/channel/services/message/ThreadMessageActivity';
import {assertThreadAllowed} from '@app/api/channel/services/thread/ThreadDenials';
import {enqueueThreadSearchSync} from '@app/api/channel/threads/ThreadJobs';
import {SYSTEM_USER_ID} from '@app/api/constants/Core';
import type {MessageAttachment, MessageReference} from '@app/api/database/types/MessageTypes';
import type {ThreadViewer} from '@app/api/experiment/ChannelThreadsGate';
import type {IFavoriteMemeRepository} from '@app/api/favorite_meme/IFavoriteMemeRepository';
import type {GatewayChannelMention, IGatewayService} from '@app/api/infrastructure/IGatewayService';
import type {ISnowflakeService} from '@app/api/infrastructure/ISnowflakeService';
@@ -63,6 +68,7 @@ import {
SENDABLE_MESSAGE_FLAGS,
} from '@fluxer/constants/src/ChannelConstants';
import {GuildNSFWLevel, GuildOperations} from '@fluxer/constants/src/GuildConstants';
import {threadWriteBlock} from '@fluxer/constants/src/ThreadPermissionUtils';
import {
DELETED_USER_ID,
RelationshipTypes,
@@ -83,6 +89,7 @@ import type {GuildResponse} from '@fluxer/schema/src/domains/guild/GuildResponse
import type {IRateLimitService} from '@pkgs/rate_limit/src/IRateLimitService';
interface MessageSendServiceDeps {
threadActivity: ThreadMessageActivity;
channelRepository: IChannelRepositoryAggregate;
userRepository: IUserRepository;
storageService: IStorageService;
@@ -303,24 +310,120 @@ export class MessageSendService {
return {canEmbedLinks, canMentionEveryone, canAttachFiles};
}
private assertThreadSendAllowed(authChannel: AuthenticatedChannel): void {
if (authChannel.thread) {
assertThreadAllowed(threadWriteBlock('send', authChannel.thread.actor));
}
}
private assertForwardableReference(isForwardMessage: boolean, referencedMessage: Message | null): void {
if (isForwardMessage && referencedMessage?.type === MessageTypes.THREAD_CREATED) {
throw InputValidationError.fromCode('message_reference', ValidationErrorCodes.CANNOT_REPLY_TO_SYSTEM_MESSAGE);
}
}
private async generateMessageId(channel: Channel): Promise<MessageID> {
let messageId = createMessageID(await this.deps.snowflakeService.generateForChannel(channel.id));
for (let attempt = 0; channel.isThread() && BigInt(messageId) <= BigInt(channel.id) && attempt < 3; attempt++) {
messageId = createMessageID(await this.deps.snowflakeService.generateForChannel(channel.id));
}
if (channel.isThread() && BigInt(messageId) <= BigInt(channel.id)) {
throw new Error('Thread message id must be greater than the thread id');
}
return messageId;
}
async validateForumStarter({
user,
parentAuth,
data,
}: {
user: User;
parentAuth: AuthenticatedChannel;
data: MessageRequest;
}): Promise<void> {
if (!user.isBot && user.id !== SYSTEM_USER_ID && !(user.flags & UserFlags.HAS_SESSION_STARTED)) {
throw InputValidationError.fromCode('content', ValidationErrorCodes.MUST_START_SESSION_BEFORE_SENDING);
}
const {channel, guild, member, checkPermission, hasPermission} = parentAuth;
await this.checkMessageSendPermissions({guild, member, channel, data, user, checkPermission, hasPermission});
this.ensureMessageRequestIsValid({user, data, guildFeatures: guild?.features ?? null});
this.deps.embedAttachmentResolver.validateAttachmentReferences({
embeds: data.embeds,
attachments: data.attachments,
});
await this.ensureAttachmentsExist({
attachments: data.attachments,
user,
channelId: channel.id,
guildFeatures: guild?.features ?? null,
});
}
async validateWebhookForumStarter({parent, data}: {parent: Channel; data: MessageRequest}): Promise<void> {
if (!parent.guildId) throw new CannotExecuteOnDmError();
const guild = await this.deps.gatewayService.getGuildData({
guildId: parent.guildId,
userId: createUserID(0n),
skipMembershipCheck: true,
});
this.ensureWebhookMessageRequestIsValid(data, guild?.features ?? null);
this.deps.embedAttachmentResolver.validateAttachmentReferences({
embeds: data.embeds,
attachments: data.attachments,
});
}
private ensureWebhookMessageRequestIsValid(data: MessageRequest, guildFeatures: Iterable<string> | null): boolean {
const isForwardMessage = data.message_reference?.type === MessageReferenceTypes.FORWARD;
if (isForwardMessage) {
if (!data.message_reference?.channel_id || !data.message_reference?.message_id) {
throw InputValidationError.fromCode(
'message_reference',
ValidationErrorCodes.FORWARD_REFERENCE_REQUIRES_CHANNEL_AND_MESSAGE,
);
}
if (
data.content ||
(data.embeds && data.embeds.length > 0) ||
(data.attachments && data.attachments.length > 0) ||
(data.sticker_ids && data.sticker_ids.length > 0)
) {
throw InputValidationError.fromCode(
'message_reference',
ValidationErrorCodes.FORWARD_MESSAGES_CANNOT_CONTAIN_CONTENT,
);
}
} else {
this.deps.validationService.validateMessageContent(data, null, {
guildFeatures,
messageAuthorType: 'webhook',
});
}
return isForwardMessage;
}
async validateMessageCanBeSent({
user,
viewer,
channelId,
data,
}: {
user: User;
viewer: ThreadViewer;
channelId: ChannelID;
data: MessageRequest;
}): Promise<void> {
const authChannel = await this.deps.channelAuthService.getChannelAuthenticated({
userId: user.id,
viewer,
channelId,
});
if (!user.isBot && user.id !== SYSTEM_USER_ID && !(user.flags & UserFlags.HAS_SESSION_STARTED)) {
throw InputValidationError.fromCode('content', ValidationErrorCodes.MUST_START_SESSION_BEFORE_SENDING);
}
if (isPersonalNotesChannel({userId: user.id, channelId})) {
await this.validatePersonalNoteMessage({user, channelId, data});
await this.validatePersonalNoteMessage({user, viewer, channelId, data});
return;
}
const {channel, guild, checkPermission, hasPermission, member} = authChannel;
@@ -334,6 +437,7 @@ export class MessageSendService {
hasPermission,
});
this.deps.validationService.ensureTextChannel(channel);
this.assertThreadSendAllowed(authChannel);
const isForwardMessage = this.ensureMessageRequestIsValid({user, data, guildFeatures: guild?.features ?? null});
this.deps.embedAttachmentResolver.validateAttachmentReferences({
embeds: data.embeds,
@@ -342,9 +446,12 @@ export class MessageSendService {
const {referencedMessage, referencedChannelGuildId} = await this.fetchReferencedMessageForValidation({
data,
channelId,
channelIsThread: channel.isThread(),
isForwardMessage,
user,
viewer,
});
this.assertForwardableReference(isForwardMessage, referencedMessage);
if (isForwardMessage && referencedMessage && guild) {
const hasEmbeds =
(referencedMessage.flags & MessageFlags.SUPPRESS_EMBEDS) === 0 && referencedMessage.embeds.length > 0;
@@ -409,15 +516,18 @@ export class MessageSendService {
private async validatePersonalNoteMessage({
user,
viewer,
channelId,
data,
}: {
user: User;
viewer: ThreadViewer;
channelId: ChannelID;
data: MessageRequest;
}): Promise<void> {
const authChannel = await this.deps.channelAuthService.getChannelAuthenticated({
userId: user.id,
viewer,
channelId,
});
const {channel} = authChannel;
@@ -430,8 +540,10 @@ export class MessageSendService {
const {referencedMessage, referencedChannelGuildId} = await this.fetchReferencedMessageForValidation({
data,
channelId,
channelIsThread: channel.isThread(),
isForwardMessage,
user,
viewer,
});
if (data.message_reference && referencedMessage && !isForwardMessage) {
const replyableTypes: ReadonlySet<Message['type']> = new Set([MessageTypes.DEFAULT, MessageTypes.REPLY]);
@@ -477,13 +589,17 @@ export class MessageSendService {
private async fetchReferencedMessageForValidation({
data,
channelId,
channelIsThread,
isForwardMessage,
user,
viewer,
}: {
data: MessageRequest;
channelId: ChannelID;
channelIsThread: boolean;
isForwardMessage: boolean;
user: User;
viewer: ThreadViewer;
}): Promise<{
referencedMessage: Message | null;
referencedChannelGuildId?: GuildID | null;
@@ -497,6 +613,7 @@ export class MessageSendService {
if (isForwardMessage) {
forwardReferenceAuthChannel = await this.deps.channelAuthService.getChannelAuthenticated({
userId: user.id,
viewer,
channelId: createChannelID(data.message_reference.channel_id!),
});
await this.ensureForwardSourceAccess(forwardReferenceAuthChannel);
@@ -508,11 +625,26 @@ export class MessageSendService {
createMessageID(data.message_reference.message_id),
);
if (!referencedMessage) {
this.assertNotThreadStarterReference(
forwardReferenceAuthChannel?.channel.isThread() ?? channelIsThread,
referenceChannelId,
data.message_reference.message_id,
);
throw new UnknownMessageError();
}
return {referencedMessage, referencedChannelGuildId};
}
private assertNotThreadStarterReference(
referenceIsThread: boolean,
referenceChannelId: ChannelID,
messageId: bigint | string,
): void {
if (referenceIsThread && String(messageId) === referenceChannelId.toString()) {
throw InputValidationError.fromCode('message_reference', ValidationErrorCodes.CANNOT_REPLY_TO_SYSTEM_MESSAGE);
}
}
private async ensureAttachmentsExist({
attachments,
user,
@@ -598,13 +730,17 @@ export class MessageSendService {
private async resolveReferenceContext({
data,
channelId,
channelIsThread,
isForwardMessage,
user,
viewer,
}: {
data: MessageRequest;
channelId: ChannelID;
channelIsThread: boolean;
isForwardMessage: boolean;
user: User;
viewer: ThreadViewer;
}): Promise<{
referencedMessage: Message | null;
referencedChannelGuildId?: GuildID | null;
@@ -616,6 +752,7 @@ export class MessageSendService {
if (isForwardMessage) {
forwardReferenceAuthChannel = await this.deps.channelAuthService.getChannelAuthenticated({
userId: user.id,
viewer,
channelId: createChannelID(data.message_reference!.channel_id!),
});
await this.ensureForwardSourceAccess(forwardReferenceAuthChannel);
@@ -629,6 +766,11 @@ export class MessageSendService {
)
: null;
if (data.message_reference && !referencedMessage) {
this.assertNotThreadStarterReference(
forwardReferenceAuthChannel?.channel.isThread() ?? channelIsThread,
referenceChannelId,
data.message_reference.message_id,
);
throw new UnknownMessageError();
}
let messageSnapshots: Array<MessageSnapshot> | undefined;
@@ -770,28 +912,35 @@ export class MessageSendService {
async sendMessage({
user,
viewer,
channelId,
data,
requestCache,
forumStarter,
}: {
user: User;
viewer: ThreadViewer;
channelId: ChannelID;
data: MessageRequest;
requestCache: RequestCache;
forumStarter?: {parentAuth: AuthenticatedChannel};
}): Promise<SendMessageResult> {
const authChannel = await this.deps.channelAuthService.getChannelAuthenticated({
userId: user.id,
viewer,
channelId,
});
if (!user.isBot && user.id !== SYSTEM_USER_ID && !(user.flags & UserFlags.HAS_SESSION_STARTED)) {
throw InputValidationError.fromCode('content', ValidationErrorCodes.MUST_START_SESSION_BEFORE_SENDING);
}
if (isPersonalNotesChannel({userId: user.id, channelId})) {
const message = await this.sendPersonalNoteMessage({authChannel, user, channelId, data, requestCache});
const message = await this.sendPersonalNoteMessage({authChannel, user, viewer, channelId, data, requestCache});
return {message, authChannel};
}
assertAccountNotLimited(user);
const {channel, guild, checkPermission, hasPermission, member} = authChannel;
const {channel, guild, member} = authChannel;
const {checkPermission, hasPermission} = forumStarter?.parentAuth ?? authChannel;
const uploadChannelId = forumStarter?.parentAuth.channel.id ?? channelId;
const {canEmbedLinks, canMentionEveryone, canAttachFiles} = await this.checkMessageSendPermissions({
guild,
member,
@@ -801,10 +950,12 @@ export class MessageSendService {
checkPermission,
hasPermission,
});
const needsSlowmodeCheck = guild && channel.rateLimitPerUser && channel.rateLimitPerUser > 0 && !user.isBot;
const needsSlowmodeCheck =
!forumStarter && guild && channel.rateLimitPerUser && channel.rateLimitPerUser > 0 && !user.isBot;
const slowmodeBypass = needsSlowmodeCheck ? await hasPermission(Permissions.BYPASS_SLOWMODE) : false;
const slowmodeKey = needsSlowmodeCheck && !slowmodeBypass ? `slowmode:${channelId}:${user.id}` : null;
this.deps.validationService.ensureTextChannel(channel);
this.assertThreadSendAllowed(authChannel);
const isForwardMessage = this.ensureMessageRequestIsValid({user, data, guildFeatures: guild?.features ?? null});
this.deps.embedAttachmentResolver.validateAttachmentReferences({
embeds: data.embeds,
@@ -821,10 +972,13 @@ export class MessageSendService {
const referenceContext = await this.resolveReferenceContext({
data,
channelId,
channelIsThread: channel.isThread(),
isForwardMessage,
user,
viewer,
});
const {referencedMessage, referencedChannelGuildId, messageSnapshots} = referenceContext;
this.assertForwardableReference(isForwardMessage, referencedMessage);
if (isForwardMessage && messageSnapshots && guild) {
const snapshotHasEmbeds = messageSnapshots.some((s) => s.embeds.length > 0);
const snapshotHasAttachments = messageSnapshots.some((s) => s.attachments.length > 0);
@@ -837,7 +991,7 @@ export class MessageSendService {
}
if (isForwardMessage && messageSnapshots && this.snapshotsContainNsfwContent(messageSnapshots)) {
const guildNsfw = guild != null && guild.nsfw_level === GuildNSFWLevel.AGE_RESTRICTED;
const destAllowsNsfw = channel.isNsfw || guildNsfw;
const destAllowsNsfw = (authChannel.thread?.parent ?? channel).isNsfw || guildNsfw;
if (!destAllowsNsfw) {
throw new NsfwContentRequiresAgeVerificationError();
}
@@ -871,7 +1025,7 @@ export class MessageSendService {
await this.ensureAttachmentsExist({
attachments: data.attachments,
user,
channelId,
channelId: uploadChannelId,
guildFeatures: guild?.features ?? null,
});
const {attachmentsToProcess, favoriteMemeAttachment} = await this.prepareMessageAttachments({
@@ -880,7 +1034,7 @@ export class MessageSendService {
data,
});
const dmNsfwContext = guild ? undefined : await this.buildDmNsfwContext(channel, user.id);
const messageId = createMessageID(await this.deps.snowflakeService.generateForChannel(channelId));
const messageId = forumStarter ? channelIdToMessageId(channel.id) : await this.generateMessageId(channel);
let mentionData: SendMentionData | undefined;
const shouldExtractMentions =
channel && !isForwardMessage && (data.content !== undefined || data.message_reference != null);
@@ -941,6 +1095,16 @@ export class MessageSendService {
});
}
}
if (authChannel.thread && !forumStarter) {
await this.deps.threadActivity.beforeUserSend({
channel,
parent: authChannel.thread.parent,
state: authChannel.thread.state,
member: authChannel.thread.member,
userId: user.id,
isBot: user.isBot,
});
}
const suppressDmRecipientDelivery = dmRecipientId !== null && isDirectDeliverySuppressed(user);
const suppressDelivery = suppressDmRecipientDelivery || isContentHidden(user, messageId);
const channelHadMessages = channel.lastMessageId !== null;
@@ -954,6 +1118,7 @@ export class MessageSendService {
embeds: data.embeds,
attachments: attachmentsToProcess,
attachmentUploadUserId: user.id,
uploadChannelId,
processedAttachments: favoriteMemeAttachment ? [favoriteMemeAttachment] : undefined,
stickerIds: data.sticker_ids ? data.sticker_ids.flatMap((stickerId) => createStickerID(stickerId)) : undefined,
messageReference,
@@ -968,6 +1133,7 @@ export class MessageSendService {
mentionData,
allowEmbeds: canEmbedLinks,
dmNsfwContext,
threadInsert: authChannel.thread !== undefined,
});
this.cacheMentionChannels({
requestCache,
@@ -996,6 +1162,20 @@ export class MessageSendService {
step: 'update_read_states',
promise: this.deps.processingService.updateReadStates({user, guild, channel, channelId, messageId}),
},
...(authChannel.thread && mentionData && mentionData.mentionUserIds.length > 0
? [
{
step: 'thread_mention_members',
promise: this.deps.threadActivity.addMentionedUsers({
channel,
parent: authChannel.thread.parent,
isModerator: authChannel.thread.isModerator,
authorId: user.id,
mentionUserIds: mentionData.mentionUserIds,
}),
},
]
: []),
]);
}
await this.settlePostCreateWork(messageId, [
@@ -1038,6 +1218,7 @@ export class MessageSendService {
void enqueueDeferredEmbeds().catch((error) => {
Logger.warn({error, messageId: messageId.toString()}, 'Failed to enqueue deferred embed extraction');
});
if (authChannel.thread) enqueueThreadSearchSync(channel.id, {activity: true});
const searchIndexOptions = this.getSearchIndexOptions(channel);
if (searchIndexOptions && !suppressDmRecipientDelivery) {
void this.deps.searchService.indexMessage(message, user.isBot, searchIndexOptions);
@@ -1047,19 +1228,23 @@ export class MessageSendService {
async sendWebhookMessage({
webhook,
thread,
data,
username,
avatar,
requestCache,
forumStarter,
}: {
webhook: Webhook;
thread?: Channel | null;
data: MessageRequest;
username?: string | null;
avatar?: string | null;
requestCache: RequestCache;
forumStarter?: boolean;
}): Promise<Message> {
const channelId = webhook.channelId!;
const channel = await this.deps.channelRepository.channelData.findUnique(channelId);
const channelId = thread?.id ?? webhook.channelId!;
const channel = thread ?? (await this.deps.channelRepository.channelData.findUnique(channelId));
if (!channel?.guildId) {
throw new CannotExecuteOnDmError();
}
@@ -1068,31 +1253,7 @@ export class MessageSendService {
userId: createUserID(0n),
skipMembershipCheck: true,
});
const isForwardMessage = data.message_reference?.type === MessageReferenceTypes.FORWARD;
if (isForwardMessage) {
if (!data.message_reference?.channel_id || !data.message_reference?.message_id) {
throw InputValidationError.fromCode(
'message_reference',
ValidationErrorCodes.FORWARD_REFERENCE_REQUIRES_CHANNEL_AND_MESSAGE,
);
}
if (
data.content ||
(data.embeds && data.embeds.length > 0) ||
(data.attachments && data.attachments.length > 0) ||
(data.sticker_ids && data.sticker_ids.length > 0)
) {
throw InputValidationError.fromCode(
'message_reference',
ValidationErrorCodes.FORWARD_MESSAGES_CANNOT_CONTAIN_CONTENT,
);
}
} else {
this.deps.validationService.validateMessageContent(data, null, {
guildFeatures: guild?.features ?? null,
messageAuthorType: 'webhook',
});
}
const isForwardMessage = this.ensureWebhookMessageRequestIsValid(data, guild?.features ?? null);
this.deps.embedAttachmentResolver.validateAttachmentReferences({
embeds: data.embeds,
attachments: data.attachments,
@@ -1136,6 +1297,7 @@ export class MessageSendService {
throw InputValidationError.fromCode('message_reference', ValidationErrorCodes.CANNOT_REPLY_TO_SYSTEM_MESSAGE);
}
}
this.assertForwardableReference(isForwardMessage, referencedMessage);
}
const messageReference = this.buildMessageReferencePayload({
data,
@@ -1144,7 +1306,10 @@ export class MessageSendService {
isForwardMessage,
referencedChannelGuildId: channel.guildId,
});
const messageId = createMessageID(await this.deps.snowflakeService.generateForChannel(channelId));
if (thread && !forumStarter) {
await this.deps.threadActivity.beforeWebhookSend(thread);
}
const messageId = thread && forumStarter ? channelIdToMessageId(thread.id) : await this.generateMessageId(channel);
let mentionData: SendMentionData | undefined;
const shouldExtractWebhookMentions =
channel && !isForwardMessage && (data.content !== undefined || data.message_reference != null);
@@ -1194,6 +1359,7 @@ export class MessageSendService {
embeds: data.embeds,
attachments: this.attachmentsToProcess(data.attachments),
attachmentUploadUserId: await this.resolveWebhookAttachmentUploadUserId(webhook, data.attachments),
uploadChannelId: thread ? (webhook.channelId ?? undefined) : undefined,
stickerIds: data.sticker_ids ? data.sticker_ids.flatMap((stickerId) => createStickerID(stickerId)) : undefined,
messageReference,
messageSnapshots,
@@ -1203,6 +1369,7 @@ export class MessageSendService {
referencedMessage,
mentionData,
allowEmbeds: true,
threadInsert: thread != null,
});
this.cacheMentionChannels({
requestCache,
@@ -1226,6 +1393,7 @@ export class MessageSendService {
void enqueueDeferredEmbeds().catch((error) => {
Logger.warn({error, messageId: messageId.toString()}, 'Failed to enqueue deferred embed extraction');
});
if (thread) enqueueThreadSearchSync(thread.id, {activity: true});
const searchIndexOptions = this.getSearchIndexOptions(channel);
if (searchIndexOptions) {
void this.deps.searchService.indexMessage(message, false, searchIndexOptions);
@@ -1235,17 +1403,19 @@ export class MessageSendService {
async editWebhookMessage({
webhook,
thread,
messageId,
data,
requestCache,
}: {
webhook: Webhook;
thread?: Channel | null;
messageId: MessageID;
data: MessageUpdateRequest;
requestCache: RequestCache;
}): Promise<Message> {
const channelId = webhook.channelId!;
const channel = await this.deps.channelRepository.channelData.findUnique(channelId);
const channelId = thread?.id ?? webhook.channelId!;
const channel = thread ?? (await this.deps.channelRepository.channelData.findUnique(channelId));
if (!channel?.guildId) {
throw new CannotExecuteOnDmError();
}
@@ -1254,6 +1424,9 @@ export class MessageSendService {
if (existingMessage.webhookId !== webhook.id) {
throw new MissingPermissionsError();
}
if (thread) {
await this.deps.threadActivity.assertWebhookCanEdit(thread);
}
const guild = await this.deps.gatewayService.getGuildData({
guildId: channel.guildId,
userId: createUserID(0n),
@@ -1309,12 +1482,14 @@ export class MessageSendService {
private async sendPersonalNoteMessage({
authChannel,
user,
viewer,
channelId,
data,
requestCache,
}: {
authChannel: AuthenticatedChannel;
user: User;
viewer: ThreadViewer;
channelId: ChannelID;
data: MessageRequest;
requestCache: RequestCache;
@@ -1336,8 +1511,10 @@ export class MessageSendService {
const {referencedMessage, referencedChannelGuildId, messageSnapshots} = await this.resolveReferenceContext({
data,
channelId,
channelIsThread: channel.isThread(),
isForwardMessage,
user,
viewer,
});
this.ensureForwardGuildMatches({data, referencedChannelGuildId});
await this.ensureAttachmentsExist({
@@ -30,6 +30,7 @@ import {
MAX_MESSAGE_LENGTH_PREMIUM,
MAX_VOICE_MESSAGE_DURATION,
} from '@fluxer/constants/src/LimitConstants';
import {THREAD_CHANNEL_TYPES} from '@fluxer/constants/src/ThreadConstants';
import {ValidationErrorCodes} from '@fluxer/constants/src/ValidationErrorCodes';
import {CannotEditSystemMessageError} from '@fluxer/errors/src/domains/channel/CannotEditSystemMessageError';
import {CannotSendEmptyMessageError} from '@fluxer/errors/src/domains/channel/CannotSendEmptyMessageError';
@@ -46,7 +47,7 @@ export class MessageValidationService {
) {}
ensureTextChannel(channel: Channel): void {
if (!TEXT_BASED_CHANNEL_TYPES.has(channel.type)) {
if (!TEXT_BASED_CHANNEL_TYPES.has(channel.type) && !THREAD_CHANNEL_TYPES.has(channel.type)) {
throw new CannotSendMessageToNonTextChannelError();
}
}
@@ -0,0 +1,189 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import type {UserID} from '@app/api/BrandedTypes';
import type {IChannelRepositoryAggregate} from '@app/api/channel/repositories/IChannelRepositoryAggregate';
import {
dispatchThreadEvents,
type ThreadDispatchEvent,
threadMembersUpdateEvent,
threadMemberUpdateEvent,
threadUpdateEvent,
} from '@app/api/channel/services/thread/ThreadDispatch';
import type {ThreadView} from '@app/api/channel/services/thread/ThreadMappers';
import {loadThreadView} from '@app/api/channel/services/thread/ThreadViews';
import {threadMemberAutojoinTotal} from '@app/api/channel/threads/ThreadMetrics';
import type {IGatewayService} from '@app/api/infrastructure/IGatewayService';
import {getKVThreadAutoArchiveQueue} from '@app/api/middleware/ServiceSingletons';
import type {Channel} from '@app/api/models/Channel';
import type {ThreadMember} from '@app/api/models/ThreadMember';
import type {ThreadState} from '@app/api/models/ThreadState';
import type {IUserRepository} from '@app/api/user/IUserRepository';
import {addThreadMembersWithinCap, threadRecipients} from '@app/api/worker/tasks/ThreadMentionScope';
import {Permissions} from '@fluxer/constants/src/ChannelConstants';
import {
MAX_ACTIVE_THREADS_PER_GUILD,
MAX_THREAD_MEMBERS,
ThreadMemberFlags,
} from '@fluxer/constants/src/ThreadConstants';
import {MaxActiveThreadsError} from '@fluxer/errors/src/domains/channel/MaxActiveThreadsError';
import {ThreadArchivedError} from '@fluxer/errors/src/domains/channel/ThreadArchivedError';
import {ThreadLockedError} from '@fluxer/errors/src/domains/channel/ThreadLockedError';
import {UnknownChannelError} from '@fluxer/errors/src/domains/channel/UnknownChannelError';
const INTERACTED_CAS_ATTEMPTS = 3;
export class ThreadMessageActivity {
constructor(
private readonly channelRepository: IChannelRepositoryAggregate,
private readonly gatewayService: IGatewayService,
private readonly userRepository: IUserRepository,
) {}
async assertCanUnarchive(state: ThreadState): Promise<void> {
if (!state.archived) return;
if ((await this.channelRepository.threads.countActiveThreads(state.guildId)) >= MAX_ACTIVE_THREADS_PER_GUILD) {
throw new MaxActiveThreadsError(MAX_ACTIVE_THREADS_PER_GUILD);
}
}
async unarchive(channel: Channel, state: ThreadState): Promise<ThreadState> {
if (!state.archived) return state;
await this.assertCanUnarchive(state);
const transition = await this.channelRepository.threads.updateState(state.threadId, (current) =>
current.archived ? {archived: false} : null,
);
if (!transition) throw new UnknownChannelError();
if (transition.previous.archived && !transition.state.archived) {
await getKVThreadAutoArchiveQueue().schedule(transition.state, channel.lastMessageId);
}
return transition.state;
}
private allMembers(thread: Channel): Promise<Array<ThreadMember>> {
return this.channelRepository.threads.listMembers(thread.id, {limit: MAX_THREAD_MEMBERS});
}
async beforeWebhookSend(thread: Channel): Promise<void> {
const state = await this.channelRepository.threads.getState(thread.id);
if (!state || state.guildId !== thread.guildId) throw new UnknownChannelError();
if (state.locked) throw new ThreadLockedError();
const next = await this.unarchive(thread, state);
if (next !== state) {
await this.dispatch(thread, [threadUpdateEvent(await this.view(thread, next), await this.allMembers(thread))]);
}
}
async assertWebhookCanEdit(thread: Channel): Promise<void> {
const state = await this.channelRepository.threads.getState(thread.id);
if (!state) throw new UnknownChannelError();
if (state.archived) throw new ThreadArchivedError();
if (state.locked) throw new ThreadLockedError();
}
async beforeUserSend(params: {
channel: Channel;
parent: Channel;
state: ThreadState;
member: ThreadMember | null;
userId: UserID;
isBot: boolean;
}): Promise<void> {
const {channel, member, userId} = params;
const state = await this.unarchive(channel, params.state);
const unarchived = state !== params.state;
const events: Array<ThreadDispatchEvent> = [];
const updated = member === null ? null : await this.markInteracted(member);
if (unarchived) {
events.push(threadUpdateEvent(await this.view(channel, state, params.parent), await this.allMembers(channel)));
} else if (updated) {
events.push(threadMemberUpdateEvent(updated));
}
if (member === null) {
const result = params.isBot
? null
: await addThreadMembersWithinCap(this.channelRepository.threads, state, [
{userId, flags: ThreadMemberFlags.HAS_INTERACTED},
]);
if (result && result.added.length > 0) {
threadMemberAutojoinTotal.inc('source="send"');
events.push(
threadMembersUpdateEvent(await this.view(channel, result.state, params.parent), {added: result.added}),
);
}
}
await this.dispatch(channel, events);
}
private async markInteracted(loaded: ThreadMember): Promise<ThreadMember | null> {
let member: ThreadMember | null = loaded;
for (let attempt = 0; attempt < INTERACTED_CAS_ATTEMPTS; attempt++) {
if (!member || (member.flags & ThreadMemberFlags.HAS_INTERACTED) !== 0) return null;
const updated = await this.channelRepository.threads.updateMemberSettings(member, {
flags: member.flags | ThreadMemberFlags.HAS_INTERACTED,
});
if (updated) return updated;
member = await this.channelRepository.threads.getMember(loaded.threadId, loaded.userId);
}
return null;
}
async addMentionedUsers(params: {
channel: Channel;
parent: Channel;
isModerator: boolean;
authorId: UserID;
mentionUserIds: Array<UserID>;
}): Promise<void> {
const {channel, parent} = params;
const guildId = channel.guildId;
if (guildId === null) return;
const mentioned = [...new Set(params.mentionUserIds)].filter((userId) => userId !== params.authorId);
if (mentioned.length === 0) return;
const recipients = await threadRecipients(this.userRepository, guildId, mentioned);
const candidates = mentioned.filter((userId) => recipients.has(userId));
if (candidates.length === 0) return;
const state = await this.channelRepository.threads.getState(channel.id);
if (!state || state.archived) return;
if (state.isPrivate && !params.isModerator && state.invitable === false) return;
const existing = new Set(
(await this.channelRepository.threads.getMembers(channel.id, candidates)).map((member) => member.userId),
);
const visible = await Promise.all(
candidates
.filter((userId) => !existing.has(userId))
.map(async (userId) =>
(await this.gatewayService.checkPermission({
guildId,
userId,
channelId: parent.id,
permission: Permissions.VIEW_CHANNEL,
}))
? userId
: null,
),
);
const toAdd = visible.filter((userId): userId is UserID => userId !== null);
if (toAdd.length === 0) return;
const result = await addThreadMembersWithinCap(
this.channelRepository.threads,
state,
toAdd.map((userId) => ({userId, flags: 0})),
);
if (!result || result.added.length === 0) return;
threadMemberAutojoinTotal.inc('source="mention"', result.added.length);
await this.dispatch(channel, [
threadMembersUpdateEvent(await this.view(channel, result.state, parent), {added: result.added}),
]);
}
async view(channel: Channel, state: ThreadState, parent?: Channel): Promise<ThreadView> {
const parentChannel = parent ?? (await this.channelRepository.channelData.findUnique(state.parentId));
if (parentChannel) return loadThreadView(this.channelRepository, channel, state, parentChannel);
return {channel, state, stats: await this.channelRepository.threads.getStats(state.threadId), parentType: null};
}
private async dispatch(channel: Channel, events: Array<ThreadDispatchEvent>): Promise<void> {
if (events.length === 0 || channel.guildId === null) return;
await dispatchThreadEvents(this.gatewayService, channel.guildId, events);
}
}

Some files were not shown because too many files have changed in this diff Show More