Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d83740e40f
|
||
|
|
8e742426c9
|
||
|
|
29edbdba18
|
||
|
|
73c103c486
|
||
|
|
7cdb27f69c
|
||
|
|
0821169a3f
|
||
|
|
1dd9db9151
|
||
|
|
d23c11cf11
|
||
|
|
5bbb0bbde5
|
||
|
|
49062159f9
|
||
|
|
1aafecf20e
|
||
|
|
68f000ef27
|
||
|
|
6d7f26f7be
|
||
|
|
446c0cd9ca
|
||
|
|
7b395e1834
|
||
|
|
b1625bfab5
|
||
|
|
b22b1ffab5
|
||
|
|
de19074fc0
|
||
|
|
ef0d1ce144
|
||
|
|
85c71faf2d
|
||
|
|
362541e13b
|
||
|
|
920f9f9f2d
|
||
|
|
23172111e9
|
||
|
|
b680f52126
|
||
|
|
eb81fa478d
|
||
|
|
68975fcadd
|
||
|
|
0b6e292cbd
|
||
|
|
68e0bc190c
|
||
|
|
c577dc5267
|
||
|
|
abe1444e4c
|
||
|
|
603a927f12
|
||
|
|
2dc8206f6f
|
||
|
|
2b65dfe076
|
||
|
|
0013726b7f
|
||
|
|
b8f509ac63
|
||
|
|
b0fdf12808
|
||
|
|
cf46b71883
|
||
|
|
bc9f7fa539
|
||
|
|
70095f3184
|
||
|
|
0fa51d6199
|
||
|
|
ca3011011a
|
||
|
|
53b8bd3fae
|
||
|
|
093426ac61
|
||
|
|
c1f967ece9
|
||
|
|
23cc721612
|
||
|
|
a23d252f27
|
||
|
|
2cf00fde4f
|
||
|
|
371f7d2ca3
|
||
|
|
12e5e92fe4
|
||
|
|
a7b4337835
|
||
|
|
a8d3a5c421
|
||
|
|
6aa8527adc
|
||
|
|
56d3c36c7d
|
||
|
|
2d13dfb063
|
||
|
|
34ae3b3e64
|
||
|
|
51f33a8dac
|
||
|
|
ee73e36f8e
|
||
|
|
0f2f43a55a
|
||
|
|
fb92bc6766
|
@@ -0,0 +1,35 @@
|
|||||||
|
name: Audit
|
||||||
|
|
||||||
|
on:
|
||||||
|
schedule:
|
||||||
|
- cron: "0 0 * * 1"
|
||||||
|
push:
|
||||||
|
paths: [uv.lock]
|
||||||
|
pull_request:
|
||||||
|
paths: [uv.lock]
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
|
||||||
|
env:
|
||||||
|
UV_PYTHON_DOWNLOADS: never
|
||||||
|
|
||||||
|
jobs:
|
||||||
|
audit:
|
||||||
|
name: Dependencies
|
||||||
|
runs-on: logaldeveloper-archlinux-ci
|
||||||
|
steps:
|
||||||
|
- name: Checkout repository
|
||||||
|
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||||
|
|
||||||
|
- name: Cache uv packages
|
||||||
|
uses: actions/cache@27d5ce7f107fe9357f9df03efb73ab90386fccae # v5.0.5
|
||||||
|
with:
|
||||||
|
path: ~/.cache/uv
|
||||||
|
key: uv-${{ hashFiles('uv.lock') }}
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: uv sync --locked
|
||||||
|
|
||||||
|
- name: Audit dependencies with pip-audit
|
||||||
|
run: uv run pip-audit --skip-editable
|
||||||
+39
-8
@@ -4,29 +4,60 @@ on:
|
|||||||
push:
|
push:
|
||||||
tags: ["v*"]
|
tags: ["v*"]
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
# packages: write # not yet supported by Gitea
|
||||||
|
|
||||||
|
env:
|
||||||
|
UV_PYTHON_DOWNLOADS: never
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
publish:
|
publish:
|
||||||
name: Publish
|
name: Publish
|
||||||
runs-on: logaldeveloper-archlinux
|
runs-on: logaldeveloper-archlinux-ci
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout repository
|
- name: Checkout repository
|
||||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||||
with:
|
with:
|
||||||
fetch-depth: 0
|
fetch-depth: 0
|
||||||
|
|
||||||
- name: Cache uv packages
|
|
||||||
uses: actions/cache@cdf6c1fa76f9f475f3d7449005a359c84ca0f306 # v5.0.3
|
|
||||||
with:
|
|
||||||
path: ~/.cache/uv
|
|
||||||
key: uv-${{ hashFiles('uv.lock') }}
|
|
||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: uv sync
|
run: uv sync --locked --no-dev
|
||||||
|
|
||||||
|
- name: Generate package metadata
|
||||||
|
id: metadata
|
||||||
|
run: |
|
||||||
|
version="${{ gitea.ref_name }}"
|
||||||
|
version="${version#v}"
|
||||||
|
printf 'version=%s\n' "$version" | tee -a "$GITEA_OUTPUT"
|
||||||
|
|
||||||
- name: Build package
|
- name: Build package
|
||||||
run: uv build
|
run: uv build
|
||||||
|
|
||||||
|
- name: Generate SBOM
|
||||||
|
env:
|
||||||
|
SYFT_CHECK_FOR_APP_UPDATE: "false"
|
||||||
|
run: |
|
||||||
|
syft scan dir:.venv \
|
||||||
|
--override-default-catalogers python-installed-package-cataloger \
|
||||||
|
--select-catalogers=-file \
|
||||||
|
--source-name git.logal.dev/LogalDeveloper/Crabstero \
|
||||||
|
--source-version "${{ steps.metadata.outputs.version }}" \
|
||||||
|
--output syft-table \
|
||||||
|
--output cyclonedx-json=crabstero-${{ steps.metadata.outputs.version }}.cyclonedx.json
|
||||||
|
sha256sum crabstero-${{ steps.metadata.outputs.version }}.cyclonedx.json
|
||||||
|
zstd -T0 --ultra -22 \
|
||||||
|
crabstero-${{ steps.metadata.outputs.version }}.cyclonedx.json
|
||||||
|
|
||||||
- name: Publish package
|
- name: Publish package
|
||||||
run: uv publish
|
run: uv publish
|
||||||
env:
|
env:
|
||||||
UV_PUBLISH_TOKEN: ${{ secrets.PYPI_TOKEN }}
|
UV_PUBLISH_TOKEN: ${{ secrets.PYPI_TOKEN }}
|
||||||
|
|
||||||
|
- name: Upload SBOM artifact
|
||||||
|
uses: christopherhx/gitea-upload-artifact@8818363695ca2d5782c64f6453273341374767b7 # v7
|
||||||
|
with:
|
||||||
|
name: crabstero-cyclonedx-${{ steps.metadata.outputs.version }}
|
||||||
|
path: crabstero-${{ steps.metadata.outputs.version }}.cyclonedx.json.zst
|
||||||
|
if-no-files-found: error
|
||||||
|
archive: "false"
|
||||||
|
|||||||
+38
-22
@@ -2,71 +2,70 @@ name: CI
|
|||||||
|
|
||||||
on:
|
on:
|
||||||
push:
|
push:
|
||||||
branches: [master]
|
|
||||||
pull_request:
|
pull_request:
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
|
||||||
|
env:
|
||||||
|
UV_PYTHON_DOWNLOADS: never
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
formatting:
|
formatting:
|
||||||
name: Formatting
|
name: Formatting
|
||||||
runs-on: logaldeveloper-archlinux
|
runs-on: logaldeveloper-archlinux-ci
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout repository
|
- name: Checkout repository
|
||||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||||
with:
|
|
||||||
fetch-depth: 0
|
|
||||||
|
|
||||||
- name: Cache uv packages
|
- name: Cache uv packages
|
||||||
uses: actions/cache@cdf6c1fa76f9f475f3d7449005a359c84ca0f306 # v5.0.3
|
uses: actions/cache@27d5ce7f107fe9357f9df03efb73ab90386fccae # v5.0.5
|
||||||
with:
|
with:
|
||||||
path: ~/.cache/uv
|
path: ~/.cache/uv
|
||||||
key: uv-${{ hashFiles('uv.lock') }}
|
key: uv-${{ hashFiles('uv.lock') }}
|
||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: uv sync
|
run: uv sync --locked
|
||||||
|
|
||||||
- name: Check formatting with Ruff
|
- name: Check formatting with Ruff
|
||||||
run: uv run ruff format --check --diff .
|
run: uv run ruff format --check --diff .
|
||||||
|
|
||||||
linting:
|
linting:
|
||||||
name: Linting
|
name: Linting
|
||||||
runs-on: logaldeveloper-archlinux
|
runs-on: logaldeveloper-archlinux-ci
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout repository
|
- name: Checkout repository
|
||||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||||
with:
|
|
||||||
fetch-depth: 0
|
|
||||||
|
|
||||||
- name: Cache uv packages
|
- name: Cache uv packages
|
||||||
uses: actions/cache@cdf6c1fa76f9f475f3d7449005a359c84ca0f306 # v5.0.3
|
uses: actions/cache@27d5ce7f107fe9357f9df03efb73ab90386fccae # v5.0.5
|
||||||
with:
|
with:
|
||||||
path: ~/.cache/uv
|
path: ~/.cache/uv
|
||||||
key: uv-${{ hashFiles('uv.lock') }}
|
key: uv-${{ hashFiles('uv.lock') }}
|
||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: uv sync
|
run: uv sync --locked
|
||||||
|
|
||||||
- name: Check linting with Ruff
|
- name: Check linting with Ruff
|
||||||
run: uv run ruff check .
|
run: uv run ruff check .
|
||||||
|
|
||||||
tests:
|
tests:
|
||||||
name: Tests
|
name: Tests
|
||||||
runs-on: logaldeveloper-archlinux
|
runs-on: logaldeveloper-archlinux-ci
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout repository
|
- name: Checkout repository
|
||||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||||
with:
|
|
||||||
fetch-depth: 0
|
|
||||||
|
|
||||||
- name: Cache uv packages
|
- name: Cache uv packages
|
||||||
uses: actions/cache@cdf6c1fa76f9f475f3d7449005a359c84ca0f306 # v5.0.3
|
uses: actions/cache@27d5ce7f107fe9357f9df03efb73ab90386fccae # v5.0.5
|
||||||
with:
|
with:
|
||||||
path: ~/.cache/uv
|
path: ~/.cache/uv
|
||||||
key: uv-${{ hashFiles('uv.lock') }}
|
key: uv-${{ hashFiles('uv.lock') }}
|
||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: uv sync
|
run: uv sync --locked
|
||||||
|
|
||||||
- name: Run unit tests with Pytest
|
- name: Run tests with Pytest
|
||||||
run: uv run pytest -v --cov --cov-report=
|
run: uv run pytest -v --cov --cov-report=
|
||||||
|
|
||||||
- name: Report code coverage
|
- name: Report code coverage
|
||||||
@@ -74,21 +73,38 @@ jobs:
|
|||||||
|
|
||||||
type-checking:
|
type-checking:
|
||||||
name: Type Checking
|
name: Type Checking
|
||||||
runs-on: logaldeveloper-archlinux
|
runs-on: logaldeveloper-archlinux-ci
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout repository
|
- name: Checkout repository
|
||||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||||
with:
|
|
||||||
fetch-depth: 0
|
|
||||||
|
|
||||||
- name: Cache uv packages
|
- name: Cache uv packages
|
||||||
uses: actions/cache@cdf6c1fa76f9f475f3d7449005a359c84ca0f306 # v5.0.3
|
uses: actions/cache@27d5ce7f107fe9357f9df03efb73ab90386fccae # v5.0.5
|
||||||
with:
|
with:
|
||||||
path: ~/.cache/uv
|
path: ~/.cache/uv
|
||||||
key: uv-${{ hashFiles('uv.lock') }}
|
key: uv-${{ hashFiles('uv.lock') }}
|
||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: uv sync
|
run: uv sync --locked
|
||||||
|
|
||||||
- name: Check types with Mypy
|
- name: Check types with Mypy
|
||||||
run: uv run mypy .
|
run: uv run mypy .
|
||||||
|
|
||||||
|
spelling:
|
||||||
|
name: Spelling
|
||||||
|
runs-on: logaldeveloper-archlinux-ci
|
||||||
|
steps:
|
||||||
|
- name: Checkout repository
|
||||||
|
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
||||||
|
|
||||||
|
- name: Cache uv packages
|
||||||
|
uses: actions/cache@27d5ce7f107fe9357f9df03efb73ab90386fccae # v5.0.5
|
||||||
|
with:
|
||||||
|
path: ~/.cache/uv
|
||||||
|
key: uv-${{ hashFiles('uv.lock') }}
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: uv sync --locked
|
||||||
|
|
||||||
|
- name: Check spelling with codespell
|
||||||
|
run: uv run codespell
|
||||||
|
|||||||
+6
-1
@@ -1,5 +1,5 @@
|
|||||||
# Version file (auto-generated by hatch-vcs)
|
# Version file (auto-generated by hatch-vcs)
|
||||||
crabstero/_version.py
|
src/crabstero/_version.py
|
||||||
|
|
||||||
# Byte-compiled / optimized / DLL files
|
# Byte-compiled / optimized / DLL files
|
||||||
__pycache__/
|
__pycache__/
|
||||||
@@ -11,6 +11,8 @@ dist/
|
|||||||
build/
|
build/
|
||||||
*.egg-info/
|
*.egg-info/
|
||||||
*.egg
|
*.egg
|
||||||
|
/*.cyclonedx.json
|
||||||
|
/*.cyclonedx.json.zst
|
||||||
|
|
||||||
# Virtual environments
|
# Virtual environments
|
||||||
venv/
|
venv/
|
||||||
@@ -23,6 +25,9 @@ env/
|
|||||||
*.swp
|
*.swp
|
||||||
*.swo
|
*.swo
|
||||||
|
|
||||||
|
# Testing
|
||||||
|
.coverage
|
||||||
|
|
||||||
# Runtime data
|
# Runtime data
|
||||||
*.db
|
*.db
|
||||||
.env
|
.env
|
||||||
|
|||||||
@@ -0,0 +1,171 @@
|
|||||||
|
# Crabstero
|
||||||
|
|
||||||
|
The simple nonversation Discord bot.
|
||||||
|
|
||||||
|
Crabstero is a Discord bot that learns from messages in each channel and uses
|
||||||
|
Markov chains to generate responses when it is mentioned.
|
||||||
|
|
||||||
|
## Overview
|
||||||
|
|
||||||
|
Crabstero is intended for Discord servers that want a lightweight chatbot with
|
||||||
|
channel-local, randomly generated replies. It does not use a large language
|
||||||
|
model; it learns from the channels it can read and uses that history to assemble
|
||||||
|
replies when mentioned.
|
||||||
|
|
||||||
|
This repository and README are for running your own instance of Crabstero. To
|
||||||
|
invite the official hosted instance instead of running your own bot, see the
|
||||||
|
[Crabstero project page](https://logal.dev/projects/crabstero/).
|
||||||
|
|
||||||
|
## Features
|
||||||
|
|
||||||
|
- Per-channel Markov chain generation for Discord messages.
|
||||||
|
- Automatic ingestion of accessible channel history and new messages.
|
||||||
|
- Mention-triggered replies based on each channel's learned message history.
|
||||||
|
- `/pingme` user control for opting into pings from generated mentions.
|
||||||
|
- `/forgetme` user data deletion and future message ingestion opt-out.
|
||||||
|
- SQLite-backed storage for learned message data.
|
||||||
|
- Optional Prometheus metrics over TCP or a Unix socket.
|
||||||
|
- systemd notification and watchdog support when run as a service.
|
||||||
|
|
||||||
|
## Requirements
|
||||||
|
|
||||||
|
- A Linux host or virtual machine.
|
||||||
|
- [uv](https://docs.astral.sh/uv/getting-started/installation/) for Python
|
||||||
|
management and the install/run commands below.
|
||||||
|
- A Discord application with a bot token and the Message Content intent enabled.
|
||||||
|
- Discord channel permissions to view channels, read message history, and send
|
||||||
|
messages where Crabstero should operate.
|
||||||
|
|
||||||
|
## Discord setup
|
||||||
|
|
||||||
|
In the [Discord Developer Portal](https://discord.com/developers/applications),
|
||||||
|
create an application, add a bot, and enable the Message Content intent for the
|
||||||
|
bot. Crabstero reads message content and embed text to learn channel-local
|
||||||
|
chains, so it cannot ingest channel history normally without that privileged
|
||||||
|
intent.
|
||||||
|
|
||||||
|
Invite the application with these OAuth2 scopes:
|
||||||
|
|
||||||
|
| Scope | Why Crabstero uses it |
|
||||||
|
|-------|------------------------|
|
||||||
|
| `bot` | Add the bot user to the selected server. |
|
||||||
|
| `applications.commands` | Install the `/pingme` and `/forgetme` slash commands. |
|
||||||
|
|
||||||
|
Recommended bot permissions:
|
||||||
|
|
||||||
|
| Permission | Why Crabstero uses it |
|
||||||
|
|------------|------------------------|
|
||||||
|
| View Channels | See the channels it should learn from and reply in. |
|
||||||
|
| Send Messages | Send generated replies when mentioned. |
|
||||||
|
| Send Messages in Threads | Reply when mentioned inside threads. |
|
||||||
|
| Embed Links | Send occasional generated embed replies. |
|
||||||
|
| Read Message History | Ingest accessible channel history on startup. |
|
||||||
|
| Use External Emojis | Preserve learned custom emoji tokens in generated replies. |
|
||||||
|
|
||||||
|
The combined permissions integer for the recommended set is
|
||||||
|
`274878254080`.
|
||||||
|
|
||||||
|
Example self-hosted invite URL:
|
||||||
|
|
||||||
|
```text
|
||||||
|
https://discord.com/oauth2/authorize?client_id=YOUR_CLIENT_ID&scope=bot%20applications.commands&permissions=274878254080
|
||||||
|
```
|
||||||
|
|
||||||
|
Replace `YOUR_CLIENT_ID` with the application ID from the Discord Developer
|
||||||
|
Portal. Channel permission overwrites still apply, so you can limit where
|
||||||
|
Crabstero learns and replies by restricting the bot role per channel.
|
||||||
|
|
||||||
|
## Installation
|
||||||
|
|
||||||
|
Create a directory for Crabstero, create a Python 3.14 virtual environment, and
|
||||||
|
install the package from the project's Gitea package index:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
mkdir crabstero && cd crabstero
|
||||||
|
uv venv --python 3.14
|
||||||
|
uv pip install crabstero \
|
||||||
|
--index https://git.logal.dev/api/packages/LogalDeveloper/pypi/simple/ \
|
||||||
|
--default-index https://pypi.org/simple
|
||||||
|
```
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
|
||||||
|
Run Crabstero with a Discord bot token from the environment:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
TOKEN="your-discord-bot-token" uv run crabstero --database-path crabstero.db
|
||||||
|
```
|
||||||
|
|
||||||
|
On startup, Crabstero opens or creates the configured SQLite database, connects
|
||||||
|
to Discord, and begins ingesting channel history and new messages it can read.
|
||||||
|
Mention the bot in a channel to receive a generated reply once enough channel
|
||||||
|
data has been collected.
|
||||||
|
|
||||||
|
## Configuration
|
||||||
|
|
||||||
|
Crabstero is configured through CLI flags, environment variables, and systemd
|
||||||
|
credentials. CLI flags take precedence over environment defaults where both are
|
||||||
|
available.
|
||||||
|
|
||||||
|
| Setting | CLI flag | Environment variable | Fallback/default | Description |
|
||||||
|
|---------|----------|----------------------|------------------|-------------|
|
||||||
|
| Discord token | `--token` | `TOKEN` | systemd credential `token` | Required Discord bot token. |
|
||||||
|
| Database path | `--database-path`, `--database` | `DATABASE_PATH` | `crabstero.db` | SQLite database file path. |
|
||||||
|
| Ingest-only mode | `--ingest-only` | | disabled | Ingest history and real-time messages without replying. |
|
||||||
|
| Metrics listener | `--listen-metrics` | `LISTEN_METRICS` | disabled | Prometheus metrics address, either `HOST:PORT` or `unix:/path.sock`. |
|
||||||
|
| Metrics socket mode | `--metrics-unix-socket-mode` | `METRICS_UNIX_SOCKET_MODE` | unchanged | Octal file mode for Unix socket metrics listeners. |
|
||||||
|
|
||||||
|
When running under systemd, Crabstero can read the bot token from a credential
|
||||||
|
named `token`. It also sends readiness, stopping, and watchdog notifications
|
||||||
|
when the relevant systemd environment variables are present.
|
||||||
|
|
||||||
|
## Development
|
||||||
|
|
||||||
|
Clone the Git repository and use [uv](https://docs.astral.sh/uv/) to sync the
|
||||||
|
project environment from the lockfile:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
git clone https://git.logal.dev/LogalDeveloper/Crabstero
|
||||||
|
cd Crabstero
|
||||||
|
uv sync --locked
|
||||||
|
```
|
||||||
|
|
||||||
|
The development toolchain uses [Ruff](https://docs.astral.sh/ruff/) for
|
||||||
|
formatting and linting, [mypy](https://mypy.readthedocs.io/en/stable/) for type
|
||||||
|
checking, [codespell](https://github.com/codespell-project/codespell) for typo
|
||||||
|
detection, [pip-audit](https://github.com/pypa/pip-audit) for dependency
|
||||||
|
auditing, and [pytest](https://docs.pytest.org/en/stable/) with
|
||||||
|
[pytest-cov](https://pytest-cov.readthedocs.io/en/latest/) and
|
||||||
|
[Coverage.py](https://coverage.readthedocs.io/) for tests and
|
||||||
|
coverage reporting.
|
||||||
|
|
||||||
|
Run the common checks:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv run ruff format --check --diff .
|
||||||
|
uv run ruff check .
|
||||||
|
uv run mypy .
|
||||||
|
uv run codespell
|
||||||
|
uv run pip-audit --skip-editable
|
||||||
|
uv run pytest
|
||||||
|
```
|
||||||
|
|
||||||
|
Run tests with coverage reporting:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
uv run pytest -v --cov --cov-report=
|
||||||
|
uv run coverage report
|
||||||
|
```
|
||||||
|
|
||||||
|
Tests are organized under `tests/unit` and `tests/integration`. The pytest
|
||||||
|
configuration marks tests automatically based on those directories.
|
||||||
|
|
||||||
|
## Links
|
||||||
|
|
||||||
|
- Project page: <https://logal.dev/projects/crabstero/>
|
||||||
|
- Source repository: <https://git.logal.dev/LogalDeveloper/Crabstero>
|
||||||
|
|
||||||
|
## License
|
||||||
|
|
||||||
|
Crabstero is licensed under the Apache License 2.0. See
|
||||||
|
[LICENSE.txt](LICENSE.txt) for the full license text.
|
||||||
@@ -1,128 +0,0 @@
|
|||||||
# Copyright 2026 Logan Fick
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
"""
|
|
||||||
Entry point for the Crabstero Discord bot.
|
|
||||||
|
|
||||||
Provides an argparse CLI with environment variable and systemd credential fallbacks for --token and --database-path.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import asyncio
|
|
||||||
import logging
|
|
||||||
import os
|
|
||||||
import signal
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from crabstero import __version__ as crabstero_version
|
|
||||||
from crabstero.bot import Crabstero
|
|
||||||
|
|
||||||
logger = logging.getLogger("crabstero")
|
|
||||||
|
|
||||||
|
|
||||||
def _read_credential(name: str) -> str | None:
|
|
||||||
"""
|
|
||||||
Read a value from a systemd credential file.
|
|
||||||
|
|
||||||
Looks for a file named *name* inside the directory pointed to by the
|
|
||||||
``CREDENTIALS_DIRECTORY`` environment variable (set automatically by
|
|
||||||
systemd when ``LoadCredential=`` or ``SetCredential=`` is used).
|
|
||||||
|
|
||||||
:param name: Credential name to look up.
|
|
||||||
:return: The credential value, or ``None`` if unavailable.
|
|
||||||
"""
|
|
||||||
credentials_dir = os.environ.get("CREDENTIALS_DIRECTORY")
|
|
||||||
if credentials_dir is None:
|
|
||||||
return None
|
|
||||||
|
|
||||||
try:
|
|
||||||
return Path(credentials_dir, name).read_text().strip()
|
|
||||||
except OSError:
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:
|
|
||||||
"""
|
|
||||||
Parses command-line arguments with environment variable fallbacks.
|
|
||||||
|
|
||||||
:param argv: Optional argument list (defaults to sys.argv[1:]).
|
|
||||||
:return: Parsed arguments namespace.
|
|
||||||
"""
|
|
||||||
parser = argparse.ArgumentParser(
|
|
||||||
prog="crabstero",
|
|
||||||
description="Crabstero - the simple nonversation Discord bot.",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--token",
|
|
||||||
default=os.environ.get("TOKEN") or _read_credential("token"),
|
|
||||||
help="Discord bot token (default: TOKEN environment variable or systemd credential 'token').",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--database-path",
|
|
||||||
"--database",
|
|
||||||
default=os.environ.get("DATABASE_PATH", "crabstero.db"),
|
|
||||||
help='Path to the SQLite database file (default: DATABASE_PATH environment variable or "crabstero.db").',
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--ingestion-workers",
|
|
||||||
type=int,
|
|
||||||
default=int(os.environ.get("INGESTION_WORKERS", "4")),
|
|
||||||
help="Number of concurrent ingestion workers (default: INGESTION_WORKERS environment variable or 4).",
|
|
||||||
)
|
|
||||||
|
|
||||||
args = parser.parse_args(argv)
|
|
||||||
|
|
||||||
if args.token is None:
|
|
||||||
parser.error(
|
|
||||||
"a Discord bot token is required via --token, the TOKEN environment variable, or a systemd credential named 'token'"
|
|
||||||
)
|
|
||||||
|
|
||||||
return args
|
|
||||||
|
|
||||||
|
|
||||||
def main() -> None:
|
|
||||||
"""
|
|
||||||
Main entry point. Parses arguments, configures logging, and starts the bot.
|
|
||||||
"""
|
|
||||||
args = _parse_args()
|
|
||||||
|
|
||||||
logging.basicConfig(
|
|
||||||
level=logging.INFO,
|
|
||||||
format="[%(asctime)s] [%(name)s] [%(levelname)s] %(message)s",
|
|
||||||
)
|
|
||||||
|
|
||||||
logger.info("Starting Crabstero %s...", crabstero_version)
|
|
||||||
|
|
||||||
bot = Crabstero(
|
|
||||||
token=args.token,
|
|
||||||
database_path=args.database_path,
|
|
||||||
ingestion_workers=args.ingestion_workers,
|
|
||||||
)
|
|
||||||
|
|
||||||
async def _run() -> None:
|
|
||||||
async with bot:
|
|
||||||
await bot.start()
|
|
||||||
|
|
||||||
# Make SIGTERM behave like SIGINT so systemd stop triggers the same
|
|
||||||
# clean shutdown path (context manager __aexit__ -> bot.close()).
|
|
||||||
signal.signal(signal.SIGTERM, signal.default_int_handler)
|
|
||||||
|
|
||||||
try:
|
|
||||||
asyncio.run(_run())
|
|
||||||
except KeyboardInterrupt:
|
|
||||||
logger.info("Interrupted.")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,153 +0,0 @@
|
|||||||
# Copyright 2026 Logan Fick
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
"""
|
|
||||||
The simple nonversation Discord bot.
|
|
||||||
|
|
||||||
Provides the Crabstero subclass that owns the full bot lifecycle: database connection,
|
|
||||||
cog loading, ingestion worker pool, and graceful shutdown.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import logging
|
|
||||||
|
|
||||||
import discord
|
|
||||||
from discord import app_commands
|
|
||||||
from discord.ext import commands
|
|
||||||
|
|
||||||
from crabstero import __version__ as crabstero_version
|
|
||||||
from crabstero.database import Database
|
|
||||||
from crabstero.tasks.ingestion import ingest_channel
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class Crabstero(commands.Bot):
|
|
||||||
"""
|
|
||||||
Central bot subclass that owns all lifecycle state.
|
|
||||||
|
|
||||||
The database is opened in setup_hook and closed in close(). Background
|
|
||||||
ingestion is handled by a bounded queue and a fixed worker pool.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self, token: str, database_path: str, ingestion_workers: int = 4
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
Configures intents, stores configuration, and prepares ingestion queue state.
|
|
||||||
|
|
||||||
:param token: The Discord bot token.
|
|
||||||
:param database_path: The file path to the SQLite database.
|
|
||||||
:param ingestion_workers: The number of concurrent ingestion workers.
|
|
||||||
"""
|
|
||||||
intents = discord.Intents.default()
|
|
||||||
intents.guilds = True
|
|
||||||
intents.guild_messages = True
|
|
||||||
intents.message_content = True
|
|
||||||
|
|
||||||
super().__init__(
|
|
||||||
command_prefix=[],
|
|
||||||
intents=intents,
|
|
||||||
max_messages=None, # Disables the message cache.
|
|
||||||
)
|
|
||||||
|
|
||||||
self._token = token
|
|
||||||
self._database_path = database_path
|
|
||||||
self._ingestion_worker_count = ingestion_workers
|
|
||||||
self._ingestion_queue: asyncio.Queue[
|
|
||||||
discord.TextChannel | discord.VoiceChannel
|
|
||||||
] = asyncio.Queue()
|
|
||||||
self._ingestion_workers: list[asyncio.Task[None]] = []
|
|
||||||
self.db: Database
|
|
||||||
|
|
||||||
self.http.user_agent = f"DiscordBot (https://git.logal.dev/LogalDeveloper/Crabstero, {crabstero_version})"
|
|
||||||
|
|
||||||
async def setup_hook(self) -> None:
|
|
||||||
"""Opens the database, starts ingestion workers, loads all cogs, and syncs slash commands if changed."""
|
|
||||||
self.db = await Database.connect(self._database_path)
|
|
||||||
self._start_ingestion_workers()
|
|
||||||
|
|
||||||
from crabstero.listeners import interaction, message, server_events
|
|
||||||
|
|
||||||
await interaction.setup(self)
|
|
||||||
await message.setup(self)
|
|
||||||
await server_events.setup(self)
|
|
||||||
|
|
||||||
# Only sync slash commands if the registered commands differ from local definitions.
|
|
||||||
local_commands = {
|
|
||||||
cmd.name: cmd.description
|
|
||||||
for cmd in self.tree.get_commands()
|
|
||||||
if isinstance(cmd, (app_commands.Command, app_commands.Group))
|
|
||||||
}
|
|
||||||
try:
|
|
||||||
remote_commands = {
|
|
||||||
cmd.name: cmd.description for cmd in await self.tree.fetch_commands()
|
|
||||||
}
|
|
||||||
except discord.HTTPException:
|
|
||||||
remote_commands = {}
|
|
||||||
|
|
||||||
if local_commands != remote_commands:
|
|
||||||
logger.info("Slash command tree has changed, syncing with Discord.")
|
|
||||||
await self.tree.sync()
|
|
||||||
|
|
||||||
async def on_ready(self) -> None:
|
|
||||||
"""Logs that the bot has started successfully."""
|
|
||||||
logger.info("Crabstero started!")
|
|
||||||
|
|
||||||
async def start(self, token: str = "", *, reconnect: bool = True) -> None:
|
|
||||||
"""
|
|
||||||
Starts the bot using the stored token by default.
|
|
||||||
|
|
||||||
:param token: Optional token override. Falls back to the stored token if empty.
|
|
||||||
:param reconnect: Whether to automatically reconnect on disconnect.
|
|
||||||
"""
|
|
||||||
await super().start(token or self._token, reconnect=reconnect)
|
|
||||||
|
|
||||||
async def close(self) -> None:
|
|
||||||
"""Cancels ingestion workers, closes the database, and then the bot connection."""
|
|
||||||
if self.is_closed():
|
|
||||||
return
|
|
||||||
logger.info("Shutting down Crabstero...")
|
|
||||||
for worker in self._ingestion_workers:
|
|
||||||
worker.cancel()
|
|
||||||
await asyncio.gather(*self._ingestion_workers, return_exceptions=True)
|
|
||||||
self._ingestion_workers.clear()
|
|
||||||
if hasattr(self, "db"):
|
|
||||||
await self.db.close()
|
|
||||||
await super().close()
|
|
||||||
|
|
||||||
def queue_channel_for_ingestion(
|
|
||||||
self, channel: discord.TextChannel | discord.VoiceChannel
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
Enqueues a single channel for background message history ingestion.
|
|
||||||
|
|
||||||
:param channel: The channel to enqueue.
|
|
||||||
"""
|
|
||||||
self._ingestion_queue.put_nowait(channel)
|
|
||||||
|
|
||||||
def _start_ingestion_workers(self) -> None:
|
|
||||||
"""Spawns the fixed pool of ingestion worker tasks."""
|
|
||||||
for _ in range(self._ingestion_worker_count):
|
|
||||||
task = asyncio.create_task(self._ingestion_worker())
|
|
||||||
self._ingestion_workers.append(task)
|
|
||||||
|
|
||||||
async def _ingestion_worker(self) -> None:
|
|
||||||
"""Loops forever pulling channels from the ingestion queue and ingesting them."""
|
|
||||||
while True:
|
|
||||||
channel = await self._ingestion_queue.get()
|
|
||||||
try:
|
|
||||||
await ingest_channel(channel, self.db)
|
|
||||||
finally:
|
|
||||||
self._ingestion_queue.task_done()
|
|
||||||
@@ -1,326 +0,0 @@
|
|||||||
# Copyright 2026 Logan Fick
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
"""
|
|
||||||
Provides async SQLite database access for all of Crabstero's persistent storage needs.
|
|
||||||
|
|
||||||
Uses aiosqlite for native async access. All methods are async def.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import contextlib
|
|
||||||
import logging
|
|
||||||
from typing import Self
|
|
||||||
|
|
||||||
import aiosqlite
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
AUTO_COMMIT_WRITE_THRESHOLD = 100 # Commit after this many DB write operations.
|
|
||||||
AUTO_COMMIT_TIMEOUT_SECONDS = (
|
|
||||||
60.0 # Commit after this many seconds since first uncommitted write.
|
|
||||||
)
|
|
||||||
|
|
||||||
# SQL statements for creating the database schema.
|
|
||||||
_SCHEMA = """
|
|
||||||
-- Markov chain starting words.
|
|
||||||
-- Each row represents one occurrence. Duplicates represent frequency weight.
|
|
||||||
CREATE TABLE IF NOT EXISTS markov_start_words (
|
|
||||||
channel_id INTEGER NOT NULL,
|
|
||||||
word TEXT NOT NULL
|
|
||||||
);
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_start_channel ON markov_start_words(channel_id, word);
|
|
||||||
|
|
||||||
-- Markov chain word transitions.
|
|
||||||
-- Each row represents one occurrence. Duplicates represent frequency weight.
|
|
||||||
CREATE TABLE IF NOT EXISTS markov_transitions (
|
|
||||||
channel_id INTEGER NOT NULL,
|
|
||||||
word TEXT NOT NULL,
|
|
||||||
next_word TEXT NOT NULL
|
|
||||||
);
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_transitions_channel_word ON markov_transitions(channel_id, word);
|
|
||||||
|
|
||||||
-- Image URLs per channel.
|
|
||||||
CREATE TABLE IF NOT EXISTS channel_images (
|
|
||||||
channel_id INTEGER NOT NULL,
|
|
||||||
url TEXT NOT NULL
|
|
||||||
);
|
|
||||||
CREATE INDEX IF NOT EXISTS idx_images_channel ON channel_images(channel_id);
|
|
||||||
|
|
||||||
-- Flags for channels, servers, and users.
|
|
||||||
CREATE TABLE IF NOT EXISTS flags (
|
|
||||||
entity_type TEXT NOT NULL,
|
|
||||||
entity_id TEXT NOT NULL,
|
|
||||||
flag_name TEXT NOT NULL,
|
|
||||||
PRIMARY KEY (entity_type, entity_id, flag_name)
|
|
||||||
);
|
|
||||||
|
|
||||||
-- Tracks which channels have been bulk-ingested.
|
|
||||||
CREATE TABLE IF NOT EXISTS ingested_channels (
|
|
||||||
channel_id INTEGER NOT NULL PRIMARY KEY
|
|
||||||
);
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
class Database:
|
|
||||||
"""
|
|
||||||
Manages all SQLite database operations for Crabstero.
|
|
||||||
|
|
||||||
Uses aiosqlite for native async access. A single connection is held open for the lifetime of
|
|
||||||
the bot process with WAL mode enabled for concurrent read performance.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, connection: aiosqlite.Connection) -> None:
|
|
||||||
"""
|
|
||||||
Initializes the Database wrapper with an already-opened aiosqlite connection.
|
|
||||||
|
|
||||||
:param connection: An open aiosqlite connection.
|
|
||||||
"""
|
|
||||||
self._connection = connection
|
|
||||||
self._pending_writes = 0
|
|
||||||
self._flush_task: asyncio.Task[None] | None = None
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
async def connect(cls, path: str) -> Self:
|
|
||||||
"""
|
|
||||||
Opens a new SQLite database at the given path, configures it for performance, and creates
|
|
||||||
the schema if it does not already exist.
|
|
||||||
|
|
||||||
:param path: The file path to the SQLite database.
|
|
||||||
:return: A new Database instance ready for use.
|
|
||||||
"""
|
|
||||||
connection = await aiosqlite.connect(path)
|
|
||||||
|
|
||||||
# Enable WAL mode for better concurrent read performance.
|
|
||||||
await connection.execute("PRAGMA journal_mode=WAL")
|
|
||||||
|
|
||||||
# Set synchronous to NORMAL for a balance between safety and speed.
|
|
||||||
await connection.execute("PRAGMA synchronous=NORMAL")
|
|
||||||
|
|
||||||
# Create the schema tables and indexes if they do not already exist.
|
|
||||||
# executescript auto-commits, so no explicit commit is needed.
|
|
||||||
await connection.executescript(_SCHEMA)
|
|
||||||
|
|
||||||
return cls(connection)
|
|
||||||
|
|
||||||
async def commit(self) -> None:
|
|
||||||
"""Commits pending writes and resets the flush timer."""
|
|
||||||
await self._connection.commit()
|
|
||||||
self._pending_writes = 0
|
|
||||||
if self._flush_task is not None:
|
|
||||||
self._flush_task.cancel()
|
|
||||||
self._flush_task = None
|
|
||||||
|
|
||||||
async def close(self) -> None:
|
|
||||||
"""Cancels the flush timer, commits any pending writes, then closes the database connection."""
|
|
||||||
if self._flush_task is not None:
|
|
||||||
self._flush_task.cancel()
|
|
||||||
with contextlib.suppress(asyncio.CancelledError):
|
|
||||||
await self._flush_task
|
|
||||||
self._flush_task = None
|
|
||||||
if self._pending_writes > 0:
|
|
||||||
await self._connection.commit()
|
|
||||||
self._pending_writes = 0
|
|
||||||
await self._connection.close()
|
|
||||||
|
|
||||||
async def add_start_words_batch(self, rows: list[tuple[int, str]]) -> None:
|
|
||||||
"""
|
|
||||||
Inserts a batch of starting words into the markov_start_words table.
|
|
||||||
|
|
||||||
:param rows: A list of (channel_id, word) tuples to insert.
|
|
||||||
"""
|
|
||||||
await self._connection.executemany(
|
|
||||||
"INSERT INTO markov_start_words (channel_id, word) VALUES (?, ?)", rows
|
|
||||||
)
|
|
||||||
await self._maybe_commit()
|
|
||||||
|
|
||||||
async def get_random_start_word(self, channel_id: int) -> str | None:
|
|
||||||
"""
|
|
||||||
Returns a random starting word for a given channel, weighted by occurrence frequency.
|
|
||||||
|
|
||||||
:param channel_id: The Discord channel ID.
|
|
||||||
:return: A random starting word, or None if no starting words exist for this channel.
|
|
||||||
"""
|
|
||||||
async with self._connection.execute(
|
|
||||||
"SELECT word FROM markov_start_words WHERE channel_id = ? ORDER BY RANDOM() LIMIT 1",
|
|
||||||
(channel_id,),
|
|
||||||
) as cursor:
|
|
||||||
row = await cursor.fetchone()
|
|
||||||
return row[0] if row else None
|
|
||||||
|
|
||||||
async def add_transitions_batch(self, rows: list[tuple[int, str, str]]) -> None:
|
|
||||||
"""
|
|
||||||
Inserts a batch of word transitions into the markov_transitions table.
|
|
||||||
|
|
||||||
:param rows: A list of (channel_id, word, next_word) tuples to insert.
|
|
||||||
"""
|
|
||||||
await self._connection.executemany(
|
|
||||||
"INSERT INTO markov_transitions (channel_id, word, next_word) VALUES (?, ?, ?)",
|
|
||||||
rows,
|
|
||||||
)
|
|
||||||
await self._maybe_commit()
|
|
||||||
|
|
||||||
async def get_random_next_word(self, channel_id: int, word: str) -> str | None:
|
|
||||||
"""
|
|
||||||
Returns a random next word for a given word in a given channel, weighted by occurrence
|
|
||||||
frequency.
|
|
||||||
|
|
||||||
:param channel_id: The Discord channel ID.
|
|
||||||
:param word: The current word to find a transition for.
|
|
||||||
:return: A random next word, or None if no transitions exist.
|
|
||||||
"""
|
|
||||||
async with self._connection.execute(
|
|
||||||
"SELECT next_word FROM markov_transitions WHERE channel_id = ? AND word = ? ORDER BY RANDOM() LIMIT 1",
|
|
||||||
(channel_id, word),
|
|
||||||
) as cursor:
|
|
||||||
row = await cursor.fetchone()
|
|
||||||
return row[0] if row else None
|
|
||||||
|
|
||||||
async def get_random_completing_next_word(
|
|
||||||
self, channel_id: int, word: str
|
|
||||||
) -> str | None:
|
|
||||||
"""
|
|
||||||
Returns a random next word that ends a sentence for a given word in a given channel,
|
|
||||||
weighted by occurrence frequency.
|
|
||||||
|
|
||||||
A completing word is one whose last character is '.', '!', '?', or '§'.
|
|
||||||
|
|
||||||
:param channel_id: The Discord channel ID.
|
|
||||||
:param word: The current word to find a completing transition for.
|
|
||||||
:return: A random completing next word, or None if no completing transitions exist.
|
|
||||||
"""
|
|
||||||
async with self._connection.execute(
|
|
||||||
"SELECT next_word FROM markov_transitions WHERE channel_id = ? AND word = ? AND SUBSTR(next_word, -1, 1) IN ('.', '!', '?', '§') ORDER BY RANDOM() LIMIT 1",
|
|
||||||
(channel_id, word),
|
|
||||||
) as cursor:
|
|
||||||
row = await cursor.fetchone()
|
|
||||||
return row[0] if row else None
|
|
||||||
|
|
||||||
async def add_image(self, channel_id: int, url: str) -> None:
|
|
||||||
"""
|
|
||||||
Stores an image URL for a given channel.
|
|
||||||
|
|
||||||
:param channel_id: The Discord channel ID.
|
|
||||||
:param url: The image URL to store.
|
|
||||||
"""
|
|
||||||
await self._connection.execute(
|
|
||||||
"INSERT INTO channel_images (channel_id, url) VALUES (?, ?)",
|
|
||||||
(channel_id, url),
|
|
||||||
)
|
|
||||||
await self._maybe_commit()
|
|
||||||
|
|
||||||
async def get_random_image(self, channel_id: int) -> str | None:
|
|
||||||
"""
|
|
||||||
Returns a random image URL for a given channel.
|
|
||||||
|
|
||||||
:param channel_id: The Discord channel ID.
|
|
||||||
:return: A random image URL, or None if no images exist for this channel.
|
|
||||||
"""
|
|
||||||
async with self._connection.execute(
|
|
||||||
"SELECT url FROM channel_images WHERE channel_id = ? ORDER BY RANDOM() LIMIT 1",
|
|
||||||
(channel_id,),
|
|
||||||
) as cursor:
|
|
||||||
row = await cursor.fetchone()
|
|
||||||
return row[0] if row else None
|
|
||||||
|
|
||||||
async def set_flag(self, entity_type: str, entity_id: str, flag_name: str) -> None:
|
|
||||||
"""
|
|
||||||
Sets a flag on a given entity. If the flag is already set, this is a no-op.
|
|
||||||
|
|
||||||
:param entity_type: The type of entity ("channel", "server", or "user").
|
|
||||||
:param entity_id: The Discord ID of the entity.
|
|
||||||
:param flag_name: The name of the flag to set.
|
|
||||||
"""
|
|
||||||
await self._connection.execute(
|
|
||||||
"INSERT OR IGNORE INTO flags (entity_type, entity_id, flag_name) VALUES (?, ?, ?)",
|
|
||||||
(entity_type, entity_id, flag_name),
|
|
||||||
)
|
|
||||||
await self._maybe_commit()
|
|
||||||
|
|
||||||
async def clear_flag(
|
|
||||||
self, entity_type: str, entity_id: str, flag_name: str
|
|
||||||
) -> None:
|
|
||||||
"""
|
|
||||||
Clears a flag on a given entity. If the flag is not set, this is a no-op.
|
|
||||||
|
|
||||||
:param entity_type: The type of entity ("channel", "server", or "user").
|
|
||||||
:param entity_id: The Discord ID of the entity.
|
|
||||||
:param flag_name: The name of the flag to clear.
|
|
||||||
"""
|
|
||||||
await self._connection.execute(
|
|
||||||
"DELETE FROM flags WHERE entity_type = ? AND entity_id = ? AND flag_name = ?",
|
|
||||||
(entity_type, entity_id, flag_name),
|
|
||||||
)
|
|
||||||
await self._maybe_commit()
|
|
||||||
|
|
||||||
async def is_flag_set(
|
|
||||||
self, entity_type: str, entity_id: str, flag_name: str
|
|
||||||
) -> bool:
|
|
||||||
"""
|
|
||||||
Checks whether a flag is set on a given entity.
|
|
||||||
|
|
||||||
:param entity_type: The type of entity ("channel", "server", or "user").
|
|
||||||
:param entity_id: The Discord ID of the entity.
|
|
||||||
:param flag_name: The name of the flag to check.
|
|
||||||
:return: True if the flag is set, False otherwise.
|
|
||||||
"""
|
|
||||||
async with self._connection.execute(
|
|
||||||
"SELECT 1 FROM flags WHERE entity_type = ? AND entity_id = ? AND flag_name = ?",
|
|
||||||
(entity_type, entity_id, flag_name),
|
|
||||||
) as cursor:
|
|
||||||
return await cursor.fetchone() is not None
|
|
||||||
|
|
||||||
async def is_channel_ingested(self, channel_id: int) -> bool:
|
|
||||||
"""
|
|
||||||
Checks whether a channel has already been bulk-ingested.
|
|
||||||
|
|
||||||
:param channel_id: The Discord channel ID.
|
|
||||||
:return: True if the channel has been ingested, False otherwise.
|
|
||||||
"""
|
|
||||||
async with self._connection.execute(
|
|
||||||
"SELECT 1 FROM ingested_channels WHERE channel_id = ?",
|
|
||||||
(channel_id,),
|
|
||||||
) as cursor:
|
|
||||||
return await cursor.fetchone() is not None
|
|
||||||
|
|
||||||
async def mark_channel_ingested(self, channel_id: int) -> None:
|
|
||||||
"""
|
|
||||||
Marks a channel as having been bulk-ingested.
|
|
||||||
|
|
||||||
:param channel_id: The Discord channel ID.
|
|
||||||
"""
|
|
||||||
await self._connection.execute(
|
|
||||||
"INSERT OR IGNORE INTO ingested_channels (channel_id) VALUES (?)",
|
|
||||||
(channel_id,),
|
|
||||||
)
|
|
||||||
await self._maybe_commit()
|
|
||||||
|
|
||||||
async def _maybe_commit(self) -> None:
|
|
||||||
"""Tracks a pending write. Commits if the threshold is reached, otherwise starts a flush timer."""
|
|
||||||
self._pending_writes += 1
|
|
||||||
if self._pending_writes >= AUTO_COMMIT_WRITE_THRESHOLD:
|
|
||||||
await self.commit()
|
|
||||||
elif self._flush_task is None:
|
|
||||||
self._flush_task = asyncio.create_task(self._flush_after_timeout())
|
|
||||||
|
|
||||||
async def _flush_after_timeout(self) -> None:
|
|
||||||
"""Background task that commits after the timeout elapses."""
|
|
||||||
try:
|
|
||||||
await asyncio.sleep(AUTO_COMMIT_TIMEOUT_SECONDS)
|
|
||||||
if self._pending_writes > 0:
|
|
||||||
await self.commit()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
raise
|
|
||||||
except Exception:
|
|
||||||
logger.exception("Error in database flush timer.")
|
|
||||||
@@ -1,97 +0,0 @@
|
|||||||
# Copyright 2026 Logan Fick
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
"""
|
|
||||||
Handles responding to interactions.
|
|
||||||
|
|
||||||
Provides the /pingme slash command as a Cog with an app command.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import logging
|
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
|
|
||||||
import discord
|
|
||||||
from discord import app_commands
|
|
||||||
from discord.ext import commands
|
|
||||||
|
|
||||||
from crabstero import flags
|
|
||||||
from crabstero.flags import EntityType, Flag
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from crabstero.bot import Crabstero
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class InteractionCog(commands.Cog):
|
|
||||||
"""
|
|
||||||
Cog for handling slash command interactions.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, bot: Crabstero) -> None:
|
|
||||||
"""
|
|
||||||
Creates a new interaction handler cog.
|
|
||||||
|
|
||||||
:param bot: The bot instance.
|
|
||||||
"""
|
|
||||||
self.bot = bot
|
|
||||||
|
|
||||||
@app_commands.command(
|
|
||||||
name="pingme",
|
|
||||||
description="Opts into (or back out of) receiving pings for generated message which mention you.",
|
|
||||||
)
|
|
||||||
async def pingme(self, interaction: discord.Interaction) -> None:
|
|
||||||
"""
|
|
||||||
Toggles the allowPings flag for the user who ran the command.
|
|
||||||
|
|
||||||
:param interaction: The interaction event.
|
|
||||||
"""
|
|
||||||
# Wrap in try/except to always send a response even if the database fails.
|
|
||||||
try:
|
|
||||||
if await flags.is_flag_set(
|
|
||||||
self.bot.db, interaction.user, EntityType.USER, Flag.ALLOW_PINGS
|
|
||||||
):
|
|
||||||
await flags.clear_flag(
|
|
||||||
self.bot.db, interaction.user, EntityType.USER, Flag.ALLOW_PINGS
|
|
||||||
)
|
|
||||||
await interaction.response.send_message(
|
|
||||||
"I will no longer ping you for messages which mention you. If you decide to opt back in, run `/pingme` any time.",
|
|
||||||
ephemeral=True,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
await flags.set_flag(
|
|
||||||
self.bot.db, interaction.user, EntityType.USER, Flag.ALLOW_PINGS
|
|
||||||
)
|
|
||||||
await interaction.response.send_message(
|
|
||||||
"I will now ping you for messages which mention you. If you change your mind, run `/pingme` any time.",
|
|
||||||
ephemeral=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
except Exception:
|
|
||||||
logger.exception(
|
|
||||||
"An exception occurred while attempting to execute slash command responder for command 'pingme'."
|
|
||||||
)
|
|
||||||
await interaction.response.send_message(
|
|
||||||
"An error occurred while executing your command. Please try again later.",
|
|
||||||
ephemeral=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
async def setup(bot: Crabstero) -> None:
|
|
||||||
"""
|
|
||||||
Adds the InteractionCog to the bot.
|
|
||||||
|
|
||||||
:param bot: The bot instance.
|
|
||||||
"""
|
|
||||||
await bot.add_cog(InteractionCog(bot))
|
|
||||||
@@ -1,79 +0,0 @@
|
|||||||
# Copyright 2026 Logan Fick
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
"""
|
|
||||||
Handles created messages.
|
|
||||||
|
|
||||||
Responds to mentions and ingests normal text messages as a Cog with an on_message listener.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
|
|
||||||
import discord
|
|
||||||
from discord.ext import commands
|
|
||||||
|
|
||||||
from crabstero.messages import ingest_message, reply_to_message
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from crabstero.bot import Crabstero
|
|
||||||
|
|
||||||
|
|
||||||
class MessageCog(commands.Cog):
|
|
||||||
"""
|
|
||||||
Cog for handling message creation events.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, bot: Crabstero) -> None:
|
|
||||||
"""
|
|
||||||
Creates a new message creation handler cog.
|
|
||||||
|
|
||||||
:param bot: The bot instance.
|
|
||||||
"""
|
|
||||||
self.bot = bot
|
|
||||||
|
|
||||||
@commands.Cog.listener()
|
|
||||||
async def on_message(self, message: discord.Message) -> None:
|
|
||||||
"""
|
|
||||||
Responds to mentions and ingests normal text messages.
|
|
||||||
|
|
||||||
:param message: The message event.
|
|
||||||
"""
|
|
||||||
channel = message.channel
|
|
||||||
|
|
||||||
if not isinstance(
|
|
||||||
channel, (discord.TextChannel, discord.Thread, discord.VoiceChannel)
|
|
||||||
):
|
|
||||||
return
|
|
||||||
|
|
||||||
if message.author.bot or message.author == self.bot.user:
|
|
||||||
return
|
|
||||||
|
|
||||||
if self.bot.user in message.mentions:
|
|
||||||
await reply_to_message(self.bot.db, message)
|
|
||||||
return
|
|
||||||
|
|
||||||
# Only ingest non-thread messages; threads share their parent channel's chain.
|
|
||||||
if message.type == discord.MessageType.default and not isinstance(
|
|
||||||
channel, discord.Thread
|
|
||||||
):
|
|
||||||
await ingest_message(self.bot.db, message)
|
|
||||||
|
|
||||||
|
|
||||||
async def setup(bot: Crabstero) -> None:
|
|
||||||
"""
|
|
||||||
Adds the MessageCog to the bot.
|
|
||||||
|
|
||||||
:param bot: The bot instance.
|
|
||||||
"""
|
|
||||||
await bot.add_cog(MessageCog(bot))
|
|
||||||
@@ -1,152 +0,0 @@
|
|||||||
# Copyright 2026 Logan Fick
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
"""
|
|
||||||
Ingests sentences and generates new ones using a Markov chain backed by SQLite storage.
|
|
||||||
|
|
||||||
The Markov chain stores word transitions per channel, using duplicate rows to represent frequency
|
|
||||||
weight. Random selection via ORDER BY RANDOM() LIMIT 1 naturally preserves this weighting.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import re
|
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from crabstero.database import Database
|
|
||||||
|
|
||||||
# Section sign (§), used internally to mark the end of sentences that lack punctuation,
|
|
||||||
# as is common on Discord.
|
|
||||||
DEFAULT_SENTENCE_END = "\u00a7"
|
|
||||||
|
|
||||||
|
|
||||||
def is_complete_sentence(sentence: str) -> bool:
|
|
||||||
"""
|
|
||||||
Checks whether a given sentence ends with a default sentence end character, a period, an
|
|
||||||
exclamation mark, or a question mark.
|
|
||||||
|
|
||||||
:param sentence: The sentence to test.
|
|
||||||
:return: True if the sentence ends with a valid terminator, False otherwise.
|
|
||||||
"""
|
|
||||||
if not sentence:
|
|
||||||
return False
|
|
||||||
|
|
||||||
return sentence[-1] in (DEFAULT_SENTENCE_END, ".", "!", "?")
|
|
||||||
|
|
||||||
|
|
||||||
async def ingest(db: Database, channel_id: int, paragraph: str) -> None:
|
|
||||||
"""
|
|
||||||
Ingests a string potentially containing multiple smaller sentences into the Markov chain
|
|
||||||
for a given channel.
|
|
||||||
|
|
||||||
:param db: The database instance.
|
|
||||||
:param channel_id: The Discord channel ID to associate with this data.
|
|
||||||
:param paragraph: The paragraph of sentences to ingest.
|
|
||||||
"""
|
|
||||||
if not is_complete_sentence(paragraph):
|
|
||||||
paragraph += DEFAULT_SENTENCE_END
|
|
||||||
|
|
||||||
# Normalize whitespace, then split on sentence-ending punctuation followed by a space.
|
|
||||||
normalized = re.sub(r" +", " ", paragraph.strip().replace("\n", " "))
|
|
||||||
sentences = re.split(r"(?<=[.!?]) ", normalized)
|
|
||||||
|
|
||||||
for sentence in sentences:
|
|
||||||
await _ingest_sentence(db, channel_id, sentence)
|
|
||||||
|
|
||||||
|
|
||||||
async def _ingest_sentence(db: Database, channel_id: int, sentence: str) -> None:
|
|
||||||
"""
|
|
||||||
Ingests a string containing a single sentence into the Markov chain for a given channel.
|
|
||||||
|
|
||||||
:param db: The database instance.
|
|
||||||
:param channel_id: The Discord channel ID to associate with this data.
|
|
||||||
:param sentence: The sentence to ingest.
|
|
||||||
"""
|
|
||||||
if not is_complete_sentence(sentence):
|
|
||||||
sentence += DEFAULT_SENTENCE_END
|
|
||||||
|
|
||||||
# Normalize whitespace and split into individual words.
|
|
||||||
words = re.sub(r" +", " ", sentence.strip()).split(" ")
|
|
||||||
|
|
||||||
start_words: list[tuple[int, str]] = []
|
|
||||||
transitions: list[tuple[int, str, str]] = []
|
|
||||||
|
|
||||||
for i in range(len(words) - 1):
|
|
||||||
if i == 0:
|
|
||||||
start_words.append((channel_id, words[i]))
|
|
||||||
transitions.append((channel_id, words[i], words[i + 1]))
|
|
||||||
|
|
||||||
if start_words:
|
|
||||||
await db.add_start_words_batch(start_words)
|
|
||||||
if transitions:
|
|
||||||
await db.add_transitions_batch(transitions)
|
|
||||||
|
|
||||||
|
|
||||||
async def generate(
|
|
||||||
db: Database, channel_id: int, soft_limit: int = 750, hard_limit: int = 1000
|
|
||||||
) -> str:
|
|
||||||
"""
|
|
||||||
Generates a new sentence using words learned from previously ingested sentences for a given
|
|
||||||
channel.
|
|
||||||
|
|
||||||
:param db: The database instance.
|
|
||||||
:param channel_id: The Discord channel ID to generate from.
|
|
||||||
:param soft_limit: The amount of characters to try and limit sentence length around.
|
|
||||||
:param hard_limit: The amount of characters to cut off the sentence at if it gets too long.
|
|
||||||
:return: A new generated sentence.
|
|
||||||
"""
|
|
||||||
word = await db.get_random_start_word(channel_id)
|
|
||||||
|
|
||||||
# Seed the chain with a fallback sentence if the channel has no data yet.
|
|
||||||
if word is None:
|
|
||||||
await _ingest_sentence(db, channel_id, "Hello world!")
|
|
||||||
word = await db.get_random_start_word(channel_id)
|
|
||||||
if word is None:
|
|
||||||
return ""
|
|
||||||
|
|
||||||
parts: list[str] = []
|
|
||||||
parts.append(word)
|
|
||||||
current_length = len(word)
|
|
||||||
|
|
||||||
# The loop is skipped if the starting word already ends a sentence (e.g. "Yes.").
|
|
||||||
while not is_complete_sentence(word):
|
|
||||||
# Past the soft limit, prefer a sentence-ending word to wrap up.
|
|
||||||
if current_length >= soft_limit:
|
|
||||||
next_word = await db.get_random_completing_next_word(channel_id, word)
|
|
||||||
if next_word is None:
|
|
||||||
next_word = await db.get_random_next_word(channel_id, word)
|
|
||||||
else:
|
|
||||||
next_word = await db.get_random_next_word(channel_id, word)
|
|
||||||
|
|
||||||
if next_word is None:
|
|
||||||
break
|
|
||||||
|
|
||||||
word = next_word
|
|
||||||
|
|
||||||
parts.append(word)
|
|
||||||
current_length += 1 + len(word) # +1 for the joining space.
|
|
||||||
|
|
||||||
if current_length >= hard_limit:
|
|
||||||
result = " ".join(parts)[:hard_limit]
|
|
||||||
# Strip the internal sentence-end marker if it ended up at the boundary.
|
|
||||||
if result and result[-1] == DEFAULT_SENTENCE_END:
|
|
||||||
return result[:-1]
|
|
||||||
return result
|
|
||||||
|
|
||||||
result = " ".join(parts)
|
|
||||||
|
|
||||||
# Strip the internal sentence-end marker so it never appears in output.
|
|
||||||
if result and result[-1] == DEFAULT_SENTENCE_END:
|
|
||||||
return result[:-1]
|
|
||||||
|
|
||||||
return result
|
|
||||||
@@ -1,162 +0,0 @@
|
|||||||
# Copyright 2026 Logan Fick
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
"""
|
|
||||||
Assists with generating Discord messages in response to other users and ingesting raw messages.
|
|
||||||
|
|
||||||
Orchestrates Markov chain generation and ingestion in the context of Discord messages, handling
|
|
||||||
reply logic, embed generation, mention filtering, and message ingestion with flag checks.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import re
|
|
||||||
import secrets
|
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
|
|
||||||
import discord
|
|
||||||
|
|
||||||
from crabstero import flags, markov
|
|
||||||
from crabstero.flags import EntityType, Flag
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from crabstero.database import Database
|
|
||||||
|
|
||||||
# Copied from discordjs/discord-api-types:
|
|
||||||
# https://github.com/discordjs/discord-api-types/blob/7fe434114e91c80ed79f0204ae6c73047672d55d/globals.ts#L30
|
|
||||||
MENTION_PATTERN = re.compile(r"<@!?(?P<id>\d{17,20})>")
|
|
||||||
|
|
||||||
_EMBED_CHANCE_THRESHOLD = 95 # Out of 100; sends an embed ~5% of the time.
|
|
||||||
|
|
||||||
|
|
||||||
async def reply_to_message(db: Database, message: discord.Message) -> None:
|
|
||||||
"""
|
|
||||||
Sends a new message in Discord in response to a given message.
|
|
||||||
|
|
||||||
:param db: The database instance.
|
|
||||||
:param message: The message prompting the response.
|
|
||||||
"""
|
|
||||||
channel = message.channel
|
|
||||||
guild = message.guild
|
|
||||||
if guild is None:
|
|
||||||
return
|
|
||||||
|
|
||||||
if not channel.permissions_for(guild.me).send_messages:
|
|
||||||
return
|
|
||||||
|
|
||||||
if (
|
|
||||||
await flags.is_flag_set(db, channel, EntityType.CHANNEL, Flag.NO_REPLY)
|
|
||||||
or await flags.is_flag_set(db, guild, EntityType.SERVER, Flag.NO_REPLY)
|
|
||||||
or await flags.is_flag_set(db, message.author, EntityType.USER, Flag.NO_REPLY)
|
|
||||||
):
|
|
||||||
return
|
|
||||||
|
|
||||||
# Threads share their parent channel's Markov chain.
|
|
||||||
if isinstance(channel, discord.Thread):
|
|
||||||
channel_id = channel.parent_id
|
|
||||||
else:
|
|
||||||
channel_id = channel.id
|
|
||||||
|
|
||||||
body = await markov.generate(db, channel_id, 750, 1000)
|
|
||||||
|
|
||||||
embed = None
|
|
||||||
|
|
||||||
# 5% chance to include an embed, if the bot has permission.
|
|
||||||
if (
|
|
||||||
secrets.randbelow(100) >= _EMBED_CHANCE_THRESHOLD
|
|
||||||
and channel.permissions_for(guild.me).embed_links
|
|
||||||
):
|
|
||||||
embed = discord.Embed(
|
|
||||||
title=await markov.generate(db, channel_id, 200, 300),
|
|
||||||
description=await markov.generate(db, channel_id, 300, 500),
|
|
||||||
)
|
|
||||||
|
|
||||||
random_image = await db.get_random_image(channel_id)
|
|
||||||
if random_image is not None:
|
|
||||||
embed.set_image(url=random_image)
|
|
||||||
|
|
||||||
# Suppress all mentions by default; only ping users who opted in.
|
|
||||||
allowed_user_ids: list[int] = []
|
|
||||||
|
|
||||||
for match in MENTION_PATTERN.finditer(body):
|
|
||||||
user_id = int(match.group("id"))
|
|
||||||
|
|
||||||
if await flags.is_flag_set(db, user_id, EntityType.USER, Flag.ALLOW_PINGS):
|
|
||||||
allowed_user_ids.append(user_id)
|
|
||||||
|
|
||||||
allowed_mentions = discord.AllowedMentions(
|
|
||||||
everyone=False,
|
|
||||||
roles=False,
|
|
||||||
users=[discord.Object(id=uid) for uid in allowed_user_ids],
|
|
||||||
)
|
|
||||||
|
|
||||||
if embed is not None:
|
|
||||||
await message.reply(
|
|
||||||
content=body,
|
|
||||||
embed=embed,
|
|
||||||
allowed_mentions=allowed_mentions,
|
|
||||||
mention_author=False,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
await message.reply(
|
|
||||||
content=body,
|
|
||||||
allowed_mentions=allowed_mentions,
|
|
||||||
mention_author=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
async def ingest_message(db: Database, message: discord.Message) -> None:
|
|
||||||
"""
|
|
||||||
Ingests a given message into its channel's Markov chain.
|
|
||||||
|
|
||||||
:param db: The database instance.
|
|
||||||
:param message: The message to ingest.
|
|
||||||
"""
|
|
||||||
guild = message.guild
|
|
||||||
if guild is None:
|
|
||||||
return
|
|
||||||
|
|
||||||
if message.author.bot or guild.me in message.mentions:
|
|
||||||
return
|
|
||||||
|
|
||||||
if (
|
|
||||||
await flags.is_flag_set(db, message.channel, EntityType.CHANNEL, Flag.NO_INGEST)
|
|
||||||
or await flags.is_flag_set(db, guild, EntityType.SERVER, Flag.NO_INGEST)
|
|
||||||
or await flags.is_flag_set(db, message.author, EntityType.USER, Flag.NO_INGEST)
|
|
||||||
):
|
|
||||||
return
|
|
||||||
|
|
||||||
channel_id = message.channel.id
|
|
||||||
|
|
||||||
if message.content:
|
|
||||||
await markov.ingest(db, channel_id, message.content)
|
|
||||||
|
|
||||||
for embed in message.embeds:
|
|
||||||
await _ingest_embed(db, channel_id, embed)
|
|
||||||
|
|
||||||
|
|
||||||
async def _ingest_embed(db: Database, channel_id: int, embed: discord.Embed) -> None:
|
|
||||||
"""
|
|
||||||
Ingests a given embed into a given channel's Markov chain.
|
|
||||||
|
|
||||||
:param db: The database instance.
|
|
||||||
:param channel_id: The ID of the channel to use for the Markov chain.
|
|
||||||
:param embed: The embed to ingest.
|
|
||||||
"""
|
|
||||||
if embed.title:
|
|
||||||
await markov.ingest(db, channel_id, embed.title)
|
|
||||||
|
|
||||||
if embed.description:
|
|
||||||
await markov.ingest(db, channel_id, embed.description)
|
|
||||||
|
|
||||||
if embed.image and embed.image.url:
|
|
||||||
await db.add_image(channel_id, embed.image.url)
|
|
||||||
+60
-26
@@ -2,29 +2,35 @@
|
|||||||
name = "crabstero"
|
name = "crabstero"
|
||||||
dynamic = ["version"]
|
dynamic = ["version"]
|
||||||
description = "The simple nonversation Discord bot."
|
description = "The simple nonversation Discord bot."
|
||||||
license = "Apache-2.0"
|
|
||||||
requires-python = ">=3.14"
|
requires-python = ">=3.14"
|
||||||
|
license = "Apache-2.0"
|
||||||
authors = [
|
authors = [
|
||||||
{ name = "Logan Fick" },
|
{ name = "Logan Fick" },
|
||||||
]
|
]
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"discord.py>=2.6.4",
|
"discord.py>=2.7.1",
|
||||||
"aiosqlite>=0.22.1",
|
"aiosqlite>=0.22.1",
|
||||||
|
"aiohttp>=3.14.1",
|
||||||
|
"prometheus-client>=0.25.0",
|
||||||
|
"uvloop>=0.22.1",
|
||||||
]
|
]
|
||||||
|
|
||||||
[project.scripts]
|
[project.scripts]
|
||||||
crabstero = "crabstero.__main__:main"
|
crabstero = "crabstero.cli:main"
|
||||||
|
|
||||||
[project.urls]
|
[project.urls]
|
||||||
Repository = "https://git.logal.dev/LogalDeveloper/Crabstero"
|
Repository = "https://git.logal.dev/LogalDeveloper/Crabstero"
|
||||||
|
|
||||||
[dependency-groups]
|
[dependency-groups]
|
||||||
dev = [
|
dev = [
|
||||||
"pytest>=9.0.2",
|
"codespell>=2.4.2",
|
||||||
"pytest-asyncio>=1.3.0",
|
"mypy>=2.1.0",
|
||||||
"pytest-cov>=7.0.0",
|
"pip-audit>=2.10.1",
|
||||||
"mypy>=1.18.1",
|
"pytest>=9.1.1",
|
||||||
"ruff>=0.15.1",
|
"pytest-asyncio>=1.4.0",
|
||||||
|
"pytest-cov>=7.1.0",
|
||||||
|
"ruff>=0.15.20",
|
||||||
|
"simcord>=1.1.0",
|
||||||
]
|
]
|
||||||
|
|
||||||
[tool.uv]
|
[tool.uv]
|
||||||
@@ -32,54 +38,82 @@ extra-index-url = ["https://git.logal.dev/api/packages/LogalDeveloper/pypi/simpl
|
|||||||
publish-url = "https://git.logal.dev/api/packages/LogalDeveloper/pypi"
|
publish-url = "https://git.logal.dev/api/packages/LogalDeveloper/pypi"
|
||||||
|
|
||||||
[build-system]
|
[build-system]
|
||||||
requires = ["hatchling>=1.28.0", "hatch-vcs>=0.5.0"]
|
requires = ["hatchling>=1.30.1", "hatch-vcs>=0.5.0"]
|
||||||
build-backend = "hatchling.build"
|
build-backend = "hatchling.build"
|
||||||
|
|
||||||
[tool.hatch.version]
|
[tool.hatch.version]
|
||||||
source = "vcs"
|
source = "vcs"
|
||||||
|
|
||||||
[tool.hatch.build.targets.wheel]
|
[tool.hatch.build.targets.wheel]
|
||||||
packages = ["crabstero"]
|
packages = ["src/crabstero"]
|
||||||
|
|
||||||
[tool.hatch.build.hooks.vcs]
|
[tool.hatch.build.hooks.vcs]
|
||||||
version-file = "crabstero/_version.py"
|
version-file = "src/crabstero/_version.py"
|
||||||
|
|
||||||
[tool.mypy]
|
[tool.mypy]
|
||||||
python_version = "3.14"
|
python_version = "3.14"
|
||||||
strict = true
|
strict = true
|
||||||
warn_unreachable = true
|
warn_unreachable = true
|
||||||
explicit_package_bases = true
|
explicit_package_bases = true
|
||||||
|
mypy_path = "$MYPY_CONFIG_FILE_DIR/src"
|
||||||
|
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
testpaths = ["tests"]
|
|
||||||
asyncio_mode = "auto"
|
asyncio_mode = "auto"
|
||||||
|
asyncio_default_fixture_loop_scope = "function"
|
||||||
|
asyncio_default_test_loop_scope = "function"
|
||||||
|
testpaths = ["tests"]
|
||||||
|
addopts = [
|
||||||
|
"--import-mode=importlib",
|
||||||
|
"--strict-config",
|
||||||
|
"--strict-markers",
|
||||||
|
]
|
||||||
|
markers = [
|
||||||
|
"unit: fast, deterministic tests that do not cross real external boundaries",
|
||||||
|
"integration: tests that exercise real local boundaries such as SQLite, sockets, CLI wiring, Discord simulation, or service lifecycle",
|
||||||
|
]
|
||||||
|
xfail_strict = true
|
||||||
|
|
||||||
[tool.ruff]
|
[tool.ruff]
|
||||||
target-version = "py314"
|
target-version = "py314"
|
||||||
extend-exclude = ["crabstero/_version.py"]
|
extend-exclude = ["src/crabstero/_version.py"] # auto-generated by hatch-vcs
|
||||||
|
|
||||||
[tool.ruff.lint]
|
[tool.ruff.lint]
|
||||||
select = [
|
select = ["ALL"]
|
||||||
"F", # Pyflakes
|
|
||||||
"E", # pycodestyle errors
|
|
||||||
"W", # pycodestyle warnings
|
|
||||||
"I", # isort
|
|
||||||
"UP", # pyupgrade
|
|
||||||
"B", # flake8-bugbear
|
|
||||||
"SIM", # flake8-simplify
|
|
||||||
"TCH", # flake8-type-checking
|
|
||||||
"RUF", # Ruff-specific rules
|
|
||||||
]
|
|
||||||
ignore = [
|
ignore = [
|
||||||
"E501", # line length enforced by the formatter; residual violations are in strings/docstrings/comments
|
"ANN401", # Any is valid at system boundaries; mypy strict handles real issues
|
||||||
|
"C901", # McCabe complexity: noisy and not actionable
|
||||||
|
"COM812", # handled by the formatter
|
||||||
|
"D203", # incompatible with D211 (no blank line before class docstring)
|
||||||
|
"D213", # incompatible with D212 (summary on first line)
|
||||||
|
"EM", # exception message style: inline literals are fine
|
||||||
|
"PLR0911", # too many return statements: flat early-returns are clear
|
||||||
|
"PLR0912", # too many branches: inherent in parsers, validators, CLI
|
||||||
|
"PLR0913", # too many arguments: API surfaces and constructors need them
|
||||||
|
"PLR0915", # too many statements: inherent in parsers, validators, CLI
|
||||||
|
"TRY003", # inline exception messages are fine (complements EM ignore)
|
||||||
|
"TRY301", # raise inside try: guard clauses don't need helper functions
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[tool.ruff.lint.per-file-ignores]
|
||||||
|
"src/crabstero/__init__.py" = ["E402"] # version fallback is resolved before public imports
|
||||||
|
"tests/**" = [
|
||||||
|
"S101", # assert is standard for pytest
|
||||||
|
"S311", # pseudo-random generators are fine in tests
|
||||||
|
"SLF001", # tests legitimately access private members for verification
|
||||||
|
"ARG001", # unused args are normal for fixtures and handler stubs
|
||||||
|
"PLR2004", # magic values are clear in test assertions
|
||||||
|
]
|
||||||
|
|
||||||
|
[tool.codespell]
|
||||||
|
skip = "src/crabstero/_version.py,uv.lock"
|
||||||
|
|
||||||
[tool.coverage.run]
|
[tool.coverage.run]
|
||||||
source = ["crabstero"]
|
source = ["crabstero"]
|
||||||
omit = [
|
omit = [
|
||||||
"crabstero/_version.py",
|
"src/crabstero/_version.py",
|
||||||
]
|
]
|
||||||
|
|
||||||
[tool.coverage.report]
|
[tool.coverage.report]
|
||||||
show_missing = true
|
show_missing = true
|
||||||
skip_empty = true
|
skip_empty = true
|
||||||
|
fail_under = 95
|
||||||
|
|||||||
@@ -0,0 +1,20 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Package execution wrapper for ``python -m crabstero``."""
|
||||||
|
|
||||||
|
from crabstero.cli import main
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
@@ -0,0 +1,361 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""The simple nonversation Discord bot.
|
||||||
|
|
||||||
|
Provides the Crabstero subclass that owns the full bot lifecycle: database connection,
|
||||||
|
cog loading, background ingestion, and graceful shutdown.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from typing import override
|
||||||
|
|
||||||
|
import discord
|
||||||
|
from discord import app_commands
|
||||||
|
from discord.app_commands import AppCommandError, CommandInvokeError
|
||||||
|
from discord.ext import commands
|
||||||
|
|
||||||
|
from crabstero import __version__ as crabstero_version
|
||||||
|
from crabstero.cache import IngestCache
|
||||||
|
from crabstero.database import Database
|
||||||
|
from crabstero.metrics import (
|
||||||
|
DISCORD_EVENTS,
|
||||||
|
DISCORD_LATENCY,
|
||||||
|
ERRORS,
|
||||||
|
GUILD_COUNT,
|
||||||
|
INGESTION_ACTIVE,
|
||||||
|
MetricsAddress,
|
||||||
|
MetricsServer,
|
||||||
|
)
|
||||||
|
from crabstero.tasks.ingestion import ingest_channel
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
type IngestableChannel = discord.TextChannel | discord.VoiceChannel
|
||||||
|
|
||||||
|
_MAX_CONCURRENT_INGESTIONS = 4
|
||||||
|
|
||||||
|
|
||||||
|
class Crabstero(commands.Bot):
|
||||||
|
"""Central bot subclass that owns all lifecycle state.
|
||||||
|
|
||||||
|
The database is opened in setup_hook and closed in close(). Background
|
||||||
|
ingestion is handled by semaphore-bounded dynamic tasks.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
database_path: str,
|
||||||
|
*,
|
||||||
|
ingest_only: bool = False,
|
||||||
|
metrics_address: MetricsAddress | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Configure intents, store configuration, and prepare ingestion state.
|
||||||
|
|
||||||
|
:param database_path: The file path to the SQLite database.
|
||||||
|
:param ingest_only: When True, the bot only ingests data and never responds.
|
||||||
|
:param metrics_address: Optional address for the Prometheus metrics server.
|
||||||
|
"""
|
||||||
|
logger.info("Crabstero v%s - A logal.dev project", crabstero_version)
|
||||||
|
|
||||||
|
intents = discord.Intents.default()
|
||||||
|
intents.guilds = True
|
||||||
|
intents.guild_messages = True
|
||||||
|
intents.message_content = True
|
||||||
|
|
||||||
|
super().__init__(
|
||||||
|
command_prefix=[],
|
||||||
|
intents=intents,
|
||||||
|
max_messages=None, # Disables the message cache.
|
||||||
|
)
|
||||||
|
|
||||||
|
self._database_path = database_path
|
||||||
|
self._ingest_only = ingest_only
|
||||||
|
self._ingestion_semaphore = asyncio.Semaphore(_MAX_CONCURRENT_INGESTIONS)
|
||||||
|
self._ingestion_tasks: dict[int, asyncio.Task[None]] = {}
|
||||||
|
self._db: Database | None = None
|
||||||
|
self.ingest_cache = IngestCache()
|
||||||
|
|
||||||
|
self._metrics_address = metrics_address
|
||||||
|
self._metrics_server: MetricsServer | None = None
|
||||||
|
|
||||||
|
DISCORD_LATENCY.set_function(lambda: self.latency)
|
||||||
|
GUILD_COUNT.set_function(lambda: len(self.guilds))
|
||||||
|
INGESTION_ACTIVE.set_function(lambda: len(self._ingestion_tasks))
|
||||||
|
|
||||||
|
repo_url = "https://git.logal.dev/LogalDeveloper/Crabstero"
|
||||||
|
self.http.user_agent = f"DiscordBot ({repo_url}, {crabstero_version})"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def db(self) -> Database:
|
||||||
|
"""The active database connection.
|
||||||
|
|
||||||
|
:raises RuntimeError: If accessed before :meth:`setup_hook` has run.
|
||||||
|
"""
|
||||||
|
if self._db is None:
|
||||||
|
msg = "Database is not initialized"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
return self._db
|
||||||
|
|
||||||
|
@property
|
||||||
|
def ingest_only(self) -> bool:
|
||||||
|
"""Whether the bot is running in ingest-only mode."""
|
||||||
|
return self._ingest_only
|
||||||
|
|
||||||
|
@override
|
||||||
|
async def setup_hook(self) -> None:
|
||||||
|
"""Open the database, start the ingest cache, and load all cogs."""
|
||||||
|
self._db = await Database.connect(self._database_path)
|
||||||
|
self.ingest_cache.start()
|
||||||
|
|
||||||
|
if self._metrics_address is not None:
|
||||||
|
server = MetricsServer(self._metrics_address)
|
||||||
|
await server.start()
|
||||||
|
self._metrics_server = server
|
||||||
|
|
||||||
|
from crabstero.listeners import ( # noqa: PLC0415
|
||||||
|
interaction,
|
||||||
|
message,
|
||||||
|
server_events,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not self._ingest_only:
|
||||||
|
await interaction.setup(self)
|
||||||
|
|
||||||
|
await message.setup(self)
|
||||||
|
await server_events.setup(self)
|
||||||
|
|
||||||
|
@self.tree.error
|
||||||
|
async def on_app_command_error(
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
error: AppCommandError,
|
||||||
|
) -> None:
|
||||||
|
original = (
|
||||||
|
error.original if isinstance(error, CommandInvokeError) else error
|
||||||
|
)
|
||||||
|
command_name = (
|
||||||
|
interaction.command.name if interaction.command else "unknown"
|
||||||
|
)
|
||||||
|
|
||||||
|
ERRORS.labels(source="command").inc()
|
||||||
|
logger.error(
|
||||||
|
"Unhandled exception in app command '%s'.",
|
||||||
|
command_name,
|
||||||
|
exc_info=original,
|
||||||
|
)
|
||||||
|
|
||||||
|
if self._metrics_server is not None:
|
||||||
|
reply = (
|
||||||
|
"I encountered an error while processing this command."
|
||||||
|
" The developer has been notified,"
|
||||||
|
" please try again later."
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
reply = (
|
||||||
|
"I encountered an error while processing this command."
|
||||||
|
" Please try again later."
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
if interaction.response.is_done():
|
||||||
|
await interaction.followup.send(reply, ephemeral=True)
|
||||||
|
else:
|
||||||
|
await interaction.response.send_message(reply, ephemeral=True)
|
||||||
|
except discord.HTTPException:
|
||||||
|
logger.debug(
|
||||||
|
"Failed to send error response for command '%s'.",
|
||||||
|
command_name,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not self._ingest_only:
|
||||||
|
# Only sync slash commands if the registered commands
|
||||||
|
# differ from local definitions.
|
||||||
|
local_commands = {
|
||||||
|
cmd.name: cmd.description
|
||||||
|
for cmd in self.tree.get_commands()
|
||||||
|
if isinstance(cmd, (app_commands.Command, app_commands.Group))
|
||||||
|
}
|
||||||
|
try:
|
||||||
|
remote_commands = {
|
||||||
|
cmd.name: cmd.description
|
||||||
|
for cmd in await self.tree.fetch_commands()
|
||||||
|
}
|
||||||
|
except discord.HTTPException:
|
||||||
|
remote_commands = {}
|
||||||
|
|
||||||
|
if local_commands != remote_commands:
|
||||||
|
logger.info("Slash command tree has changed, syncing with Discord.")
|
||||||
|
await self.tree.sync()
|
||||||
|
|
||||||
|
@override
|
||||||
|
def dispatch(self, event: str, /, *args: object, **kwargs: object) -> None:
|
||||||
|
"""Dispatch an event, incrementing the events counter.
|
||||||
|
|
||||||
|
:param event: The event name.
|
||||||
|
"""
|
||||||
|
DISCORD_EVENTS.labels(event=event).inc()
|
||||||
|
super().dispatch(event, *args, **kwargs)
|
||||||
|
|
||||||
|
@override
|
||||||
|
async def on_error(
|
||||||
|
self,
|
||||||
|
event_method: str,
|
||||||
|
/,
|
||||||
|
*args: object,
|
||||||
|
**kwargs: object,
|
||||||
|
) -> None:
|
||||||
|
"""Increment the global error counter for event listener exceptions.
|
||||||
|
|
||||||
|
:param event_method: The name of the event that raised the exception.
|
||||||
|
"""
|
||||||
|
ERRORS.labels(source=event_method).inc()
|
||||||
|
logger.error("Unhandled exception in %s.", event_method, exc_info=True) # noqa: LOG014
|
||||||
|
|
||||||
|
async def on_ready(self) -> None:
|
||||||
|
"""Log that the bot has started successfully."""
|
||||||
|
logger.info("Crabstero started!")
|
||||||
|
|
||||||
|
@override
|
||||||
|
async def close(self) -> None:
|
||||||
|
"""Cancel ingestion tasks, close Discord, and release local resources."""
|
||||||
|
if self.is_closed():
|
||||||
|
await self._close_local_resources()
|
||||||
|
return
|
||||||
|
logger.info("Shutting down Crabstero...")
|
||||||
|
tasks = list(self._ingestion_tasks.values())
|
||||||
|
self._ingestion_tasks.clear()
|
||||||
|
for task in tasks:
|
||||||
|
task.cancel()
|
||||||
|
await asyncio.gather(*tasks, return_exceptions=True)
|
||||||
|
try:
|
||||||
|
await super().close()
|
||||||
|
finally:
|
||||||
|
await self._close_local_resources()
|
||||||
|
|
||||||
|
async def _close_local_resources(self) -> None:
|
||||||
|
"""Release local resources, attempting each cleanup even if one fails."""
|
||||||
|
errors: list[Exception] = []
|
||||||
|
|
||||||
|
try:
|
||||||
|
await self.ingest_cache.stop()
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
errors.append(exc)
|
||||||
|
|
||||||
|
if self._metrics_server is not None:
|
||||||
|
try:
|
||||||
|
await self._metrics_server.stop()
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
errors.append(exc)
|
||||||
|
else:
|
||||||
|
self._metrics_server = None
|
||||||
|
|
||||||
|
if self._db is not None:
|
||||||
|
try:
|
||||||
|
await self._db.close()
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
errors.append(exc)
|
||||||
|
else:
|
||||||
|
self._db = None
|
||||||
|
|
||||||
|
if len(errors) == 1:
|
||||||
|
raise errors[0]
|
||||||
|
if errors:
|
||||||
|
raise ExceptionGroup(
|
||||||
|
"Errors occurred while closing local resources.",
|
||||||
|
errors,
|
||||||
|
)
|
||||||
|
|
||||||
|
def queue_channel_for_ingestion(self, channel: IngestableChannel) -> None:
|
||||||
|
"""Create a background task to ingest a single channel.
|
||||||
|
|
||||||
|
Duplicate requests for a channel that is already in-flight are ignored.
|
||||||
|
Concurrency is bounded by the ingestion semaphore.
|
||||||
|
|
||||||
|
:param channel: The channel to ingest.
|
||||||
|
"""
|
||||||
|
if channel.id in self._ingestion_tasks:
|
||||||
|
return
|
||||||
|
task = asyncio.create_task(
|
||||||
|
self._ingest_one(channel),
|
||||||
|
name=f"ingest-{channel.id}",
|
||||||
|
)
|
||||||
|
self._ingestion_tasks[channel.id] = task
|
||||||
|
task.add_done_callback(lambda t: self._on_ingestion_done(channel.id, t))
|
||||||
|
|
||||||
|
async def _ingest_one(self, channel: IngestableChannel) -> None:
|
||||||
|
"""Acquire the semaphore and ingest one channel."""
|
||||||
|
async with self._ingestion_semaphore:
|
||||||
|
await ingest_channel(channel, self.db)
|
||||||
|
|
||||||
|
def _on_ingestion_done(self, channel_id: int, task: asyncio.Task[None]) -> None:
|
||||||
|
"""Clean up a finished ingestion task and log any errors."""
|
||||||
|
self._ingestion_tasks.pop(channel_id, None)
|
||||||
|
if task.cancelled():
|
||||||
|
return
|
||||||
|
exc = task.exception()
|
||||||
|
if exc is not None:
|
||||||
|
ERRORS.labels(source="ingestion").inc()
|
||||||
|
logger.error("Ingestion task failed.", exc_info=exc)
|
||||||
|
|
||||||
|
|
||||||
|
class TrackedView(discord.ui.View):
|
||||||
|
"""Base View that increments the global error counter on failures."""
|
||||||
|
|
||||||
|
@override
|
||||||
|
async def on_error(
|
||||||
|
self,
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
error: Exception,
|
||||||
|
item: discord.ui.Item[TrackedView],
|
||||||
|
/,
|
||||||
|
) -> None:
|
||||||
|
"""Increment the error counter and log the exception.
|
||||||
|
|
||||||
|
:param interaction: The interaction that led to the failure.
|
||||||
|
:param error: The exception that was raised.
|
||||||
|
:param item: The item that failed the dispatch.
|
||||||
|
"""
|
||||||
|
ERRORS.labels(source="view").inc()
|
||||||
|
logger.error(
|
||||||
|
"Unhandled exception in view %r for item %r.",
|
||||||
|
self,
|
||||||
|
item,
|
||||||
|
exc_info=error,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TrackedModal(discord.ui.Modal):
|
||||||
|
"""Base Modal that increments the global error counter on failures."""
|
||||||
|
|
||||||
|
@override
|
||||||
|
async def on_error(
|
||||||
|
self,
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
error: Exception,
|
||||||
|
item: discord.ui.Item[TrackedModal] | None = None,
|
||||||
|
/,
|
||||||
|
) -> None:
|
||||||
|
"""Increment the error counter and log the exception.
|
||||||
|
|
||||||
|
:param interaction: The interaction that led to the failure.
|
||||||
|
:param error: The exception that was raised.
|
||||||
|
:param item: Unused. Present for BaseView signature compatibility.
|
||||||
|
"""
|
||||||
|
ERRORS.labels(source="modal").inc()
|
||||||
|
logger.error(
|
||||||
|
"Unhandled exception in modal %r.",
|
||||||
|
self,
|
||||||
|
exc_info=error,
|
||||||
|
)
|
||||||
@@ -0,0 +1,122 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""TTL cache for recently ingested messages, enabling uningest on delete."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import contextlib
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class CachedMessage:
|
||||||
|
"""Data stored for a recently ingested message.
|
||||||
|
|
||||||
|
:param channel_id: The Discord channel ID the message was ingested into.
|
||||||
|
:param user_id: The Discord user ID of the message author.
|
||||||
|
:param content: The text content of the message, if any.
|
||||||
|
:param embed_texts: Titles and descriptions extracted from embeds.
|
||||||
|
:param image_urls: Image URLs extracted from embeds.
|
||||||
|
"""
|
||||||
|
|
||||||
|
channel_id: int
|
||||||
|
user_id: int
|
||||||
|
content: str | None
|
||||||
|
embed_texts: list[str]
|
||||||
|
image_urls: list[str]
|
||||||
|
|
||||||
|
|
||||||
|
class IngestCache:
|
||||||
|
"""In-memory TTL cache mapping message IDs to their ingested data.
|
||||||
|
|
||||||
|
Entries expire after ``ttl_seconds`` and are evicted by a periodic
|
||||||
|
background task.
|
||||||
|
|
||||||
|
:param ttl_seconds: How long entries remain valid, in seconds.
|
||||||
|
:param cleanup_interval_seconds: How often the background task runs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("_cleanup_interval", "_entries", "_task", "_ttl")
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
ttl_seconds: float = 120,
|
||||||
|
cleanup_interval_seconds: float = 30,
|
||||||
|
) -> None:
|
||||||
|
"""Create a new cache with the given TTL and cleanup interval.
|
||||||
|
|
||||||
|
:param ttl_seconds: How long entries remain valid, in seconds.
|
||||||
|
:param cleanup_interval_seconds: How often the background task runs.
|
||||||
|
"""
|
||||||
|
self._ttl = ttl_seconds
|
||||||
|
self._cleanup_interval = cleanup_interval_seconds
|
||||||
|
self._entries: dict[int, tuple[float, CachedMessage]] = {}
|
||||||
|
self._task: asyncio.Task[None] | None = None
|
||||||
|
|
||||||
|
def put(self, message_id: int, entry: CachedMessage) -> None:
|
||||||
|
"""Store a cache entry for a message.
|
||||||
|
|
||||||
|
:param message_id: The Discord message ID.
|
||||||
|
:param entry: The cached message data.
|
||||||
|
"""
|
||||||
|
self._entries[message_id] = (time.monotonic(), entry)
|
||||||
|
|
||||||
|
def pop(self, message_id: int) -> CachedMessage | None:
|
||||||
|
"""Remove and return a cache entry if it exists and has not expired.
|
||||||
|
|
||||||
|
:param message_id: The Discord message ID.
|
||||||
|
:return: The cached message data, or None.
|
||||||
|
"""
|
||||||
|
pair = self._entries.pop(message_id, None)
|
||||||
|
if pair is None:
|
||||||
|
return None
|
||||||
|
stored_at, entry = pair
|
||||||
|
if time.monotonic() - stored_at > self._ttl:
|
||||||
|
return None
|
||||||
|
return entry
|
||||||
|
|
||||||
|
def _cleanup(self) -> None:
|
||||||
|
"""Remove all expired entries from the cache."""
|
||||||
|
now = time.monotonic()
|
||||||
|
expired = [
|
||||||
|
mid
|
||||||
|
for mid, (stored_at, _) in self._entries.items()
|
||||||
|
if now - stored_at > self._ttl
|
||||||
|
]
|
||||||
|
for mid in expired:
|
||||||
|
del self._entries[mid]
|
||||||
|
|
||||||
|
def start(self) -> None:
|
||||||
|
"""Start the periodic background cleanup task."""
|
||||||
|
if self._task is not None:
|
||||||
|
return
|
||||||
|
self._task = asyncio.create_task(
|
||||||
|
self._cleanup_loop(),
|
||||||
|
name="ingest-cache-cleanup",
|
||||||
|
)
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
"""Stop the periodic background cleanup task and wait for it to finish."""
|
||||||
|
if self._task is not None:
|
||||||
|
self._task.cancel()
|
||||||
|
with contextlib.suppress(asyncio.CancelledError):
|
||||||
|
await self._task
|
||||||
|
self._task = None
|
||||||
|
|
||||||
|
async def _cleanup_loop(self) -> None:
|
||||||
|
"""Run cleanup on a fixed interval until cancelled."""
|
||||||
|
while True:
|
||||||
|
await asyncio.sleep(self._cleanup_interval)
|
||||||
|
self._cleanup()
|
||||||
@@ -0,0 +1,402 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Command-line interface for the Crabstero Discord bot.
|
||||||
|
|
||||||
|
Provides an argparse CLI with environment variable and systemd credential
|
||||||
|
fallbacks, plus systemd notify/watchdog integration when run as a service.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import asyncio
|
||||||
|
import contextlib
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import signal
|
||||||
|
import socket
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import TYPE_CHECKING, cast
|
||||||
|
|
||||||
|
import uvloop
|
||||||
|
|
||||||
|
from crabstero.bot import Crabstero
|
||||||
|
from crabstero.metrics import MetricsAddress, TcpMetricsAddress, UnixMetricsAddress
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import Callable, Iterator
|
||||||
|
|
||||||
|
logger = logging.getLogger("crabstero")
|
||||||
|
|
||||||
|
_KEYBOARD_INTERRUPT_EXIT_CODE = 130
|
||||||
|
_MAX_METRICS_UNIX_SOCKET_MODE = 0o777
|
||||||
|
_HANDLED_SIGNALS = (signal.SIGINT, signal.SIGTERM)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class _CliConfig:
|
||||||
|
"""Resolved command-line configuration."""
|
||||||
|
|
||||||
|
token: str
|
||||||
|
database_path: str
|
||||||
|
ingest_only: bool
|
||||||
|
metrics_address: MetricsAddress | None
|
||||||
|
|
||||||
|
|
||||||
|
def _sd_notify(state: str) -> None:
|
||||||
|
"""Send a notification to systemd, if running under a systemd service.
|
||||||
|
|
||||||
|
:param state: The systemd notification payload.
|
||||||
|
"""
|
||||||
|
addr = os.environ.get("NOTIFY_SOCKET")
|
||||||
|
if not addr:
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
with socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM) as sock:
|
||||||
|
if addr[0] == "@":
|
||||||
|
addr = "\0" + addr[1:]
|
||||||
|
sock.sendto(state.encode(), addr)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.debug("Could not send systemd notification: %s", exc)
|
||||||
|
|
||||||
|
|
||||||
|
def _watchdog_interval() -> float | None:
|
||||||
|
"""Return the systemd watchdog ping interval in seconds, if enabled."""
|
||||||
|
usec = os.environ.get("WATCHDOG_USEC")
|
||||||
|
if not usec:
|
||||||
|
return None
|
||||||
|
|
||||||
|
pid = os.environ.get("WATCHDOG_PID")
|
||||||
|
if pid is not None:
|
||||||
|
try:
|
||||||
|
watchdog_pid = int(pid)
|
||||||
|
except ValueError:
|
||||||
|
return None
|
||||||
|
if watchdog_pid != os.getpid():
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
watchdog_usec = int(usec)
|
||||||
|
except ValueError:
|
||||||
|
return None
|
||||||
|
|
||||||
|
if watchdog_usec <= 0:
|
||||||
|
return None
|
||||||
|
|
||||||
|
return watchdog_usec / 1_000_000 / 2
|
||||||
|
|
||||||
|
|
||||||
|
async def _watchdog_loop(interval: float, stop: asyncio.Event) -> None:
|
||||||
|
"""Send systemd watchdog pings until *stop* is set.
|
||||||
|
|
||||||
|
:param interval: Seconds to wait between watchdog pings.
|
||||||
|
:param stop: Event that stops the watchdog loop.
|
||||||
|
"""
|
||||||
|
while not stop.is_set():
|
||||||
|
_sd_notify("WATCHDOG=1")
|
||||||
|
with contextlib.suppress(TimeoutError):
|
||||||
|
await asyncio.wait_for(stop.wait(), timeout=interval)
|
||||||
|
|
||||||
|
|
||||||
|
def _read_credential(name: str) -> str | None:
|
||||||
|
"""Read a value from a systemd credential file.
|
||||||
|
|
||||||
|
Looks for a file named *name* inside the directory pointed to by the
|
||||||
|
``CREDENTIALS_DIRECTORY`` environment variable (set automatically by
|
||||||
|
systemd when ``LoadCredential=`` or ``SetCredential=`` is used).
|
||||||
|
|
||||||
|
:param name: Credential name to look up.
|
||||||
|
:return: The credential value, or ``None`` if unavailable.
|
||||||
|
"""
|
||||||
|
credentials_dir = os.environ.get("CREDENTIALS_DIRECTORY")
|
||||||
|
if credentials_dir is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
return Path(credentials_dir, name).read_text().strip()
|
||||||
|
except OSError:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def _install_signal_handlers(
|
||||||
|
loop: asyncio.AbstractEventLoop,
|
||||||
|
callback: Callable[[signal.Signals], None],
|
||||||
|
) -> Iterator[None]:
|
||||||
|
"""Install process signal handlers and always remove installed handlers."""
|
||||||
|
registered_signals: list[signal.Signals] = []
|
||||||
|
try:
|
||||||
|
for sig in _HANDLED_SIGNALS:
|
||||||
|
loop.add_signal_handler(sig, callback, sig)
|
||||||
|
registered_signals.append(sig)
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
for sig in registered_signals:
|
||||||
|
loop.remove_signal_handler(sig)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_args(argv: list[str] | None = None) -> _CliConfig:
|
||||||
|
"""Parse command-line arguments with environment variable fallbacks.
|
||||||
|
|
||||||
|
:param argv: Optional argument list (defaults to sys.argv[1:]).
|
||||||
|
:return: Resolved command-line configuration.
|
||||||
|
"""
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
prog="crabstero",
|
||||||
|
description="Crabstero - the simple nonversation Discord bot.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--token",
|
||||||
|
default=os.environ.get("TOKEN") or _read_credential("token"),
|
||||||
|
help=(
|
||||||
|
"Discord bot token"
|
||||||
|
" (default: TOKEN environment variable"
|
||||||
|
" or systemd credential 'token')."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--database-path",
|
||||||
|
"--database",
|
||||||
|
default=os.environ.get("DATABASE_PATH", "crabstero.db"),
|
||||||
|
help=(
|
||||||
|
"Path to the SQLite database file"
|
||||||
|
" (default: DATABASE_PATH environment"
|
||||||
|
' variable or "crabstero.db").'
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--ingest-only",
|
||||||
|
action="store_true",
|
||||||
|
default=False,
|
||||||
|
help=(
|
||||||
|
"Run in ingest-only mode: ingest channel"
|
||||||
|
" history and real-time messages but"
|
||||||
|
" never respond."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--listen-metrics",
|
||||||
|
default=os.environ.get("LISTEN_METRICS"),
|
||||||
|
help=(
|
||||||
|
"Enable Prometheus metrics endpoint on HOST:PORT or unix:/path.sock"
|
||||||
|
" (e.g. 127.0.0.1:9090 or unix:/run/crabstero.sock)."
|
||||||
|
" Disabled by default."
|
||||||
|
" (default: LISTEN_METRICS environment variable)."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--metrics-unix-socket-mode",
|
||||||
|
default=os.environ.get("METRICS_UNIX_SOCKET_MODE"),
|
||||||
|
help=(
|
||||||
|
"Octal file mode to set on a metrics Unix socket"
|
||||||
|
" (e.g. 0666). Only valid with --listen-metrics unix:/path.sock."
|
||||||
|
" (default: METRICS_UNIX_SOCKET_MODE environment variable)."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args(argv)
|
||||||
|
|
||||||
|
token = cast("str | None", args.token)
|
||||||
|
database_path = cast("str", args.database_path)
|
||||||
|
ingest_only = cast("bool", args.ingest_only)
|
||||||
|
raw_metrics_address = cast("str | None", args.listen_metrics)
|
||||||
|
raw_metrics_unix_socket_mode = cast(
|
||||||
|
"str | None",
|
||||||
|
args.metrics_unix_socket_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
metrics_unix_socket_mode: int | None = None
|
||||||
|
if raw_metrics_unix_socket_mode is not None:
|
||||||
|
try:
|
||||||
|
metrics_unix_socket_mode = int(raw_metrics_unix_socket_mode, 8)
|
||||||
|
except ValueError:
|
||||||
|
parser.error(
|
||||||
|
"--metrics-unix-socket-mode must be an octal mode"
|
||||||
|
" between 0000 and 0777",
|
||||||
|
)
|
||||||
|
if not 0 <= metrics_unix_socket_mode <= _MAX_METRICS_UNIX_SOCKET_MODE:
|
||||||
|
parser.error(
|
||||||
|
"--metrics-unix-socket-mode must be an octal mode"
|
||||||
|
" between 0000 and 0777",
|
||||||
|
)
|
||||||
|
|
||||||
|
metrics_address: MetricsAddress | None = None
|
||||||
|
if raw_metrics_address is not None:
|
||||||
|
if raw_metrics_address.startswith("unix:"):
|
||||||
|
path = raw_metrics_address.removeprefix("unix:")
|
||||||
|
if not path:
|
||||||
|
parser.error("--listen-metrics Unix socket path must not be empty")
|
||||||
|
metrics_address = UnixMetricsAddress(path, metrics_unix_socket_mode)
|
||||||
|
else:
|
||||||
|
if metrics_unix_socket_mode is not None:
|
||||||
|
parser.error(
|
||||||
|
"--metrics-unix-socket-mode is only valid with"
|
||||||
|
" --listen-metrics unix:/path.sock",
|
||||||
|
)
|
||||||
|
host, sep, port_str = raw_metrics_address.rpartition(":")
|
||||||
|
if not sep or not host:
|
||||||
|
parser.error(
|
||||||
|
"--listen-metrics must be in HOST:PORT or unix:/path.sock format "
|
||||||
|
"(e.g. 127.0.0.1:9090)",
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
metrics_address = TcpMetricsAddress(host, int(port_str))
|
||||||
|
except ValueError:
|
||||||
|
parser.error(
|
||||||
|
f"--listen-metrics port must be an integer, got '{port_str}'",
|
||||||
|
)
|
||||||
|
elif metrics_unix_socket_mode is not None:
|
||||||
|
parser.error(
|
||||||
|
"--metrics-unix-socket-mode is only valid with"
|
||||||
|
" --listen-metrics unix:/path.sock",
|
||||||
|
)
|
||||||
|
|
||||||
|
if token is None:
|
||||||
|
parser.error(
|
||||||
|
"a Discord bot token is required via --token,"
|
||||||
|
" the TOKEN environment variable,"
|
||||||
|
" or a systemd credential named 'token'",
|
||||||
|
)
|
||||||
|
|
||||||
|
return _CliConfig(
|
||||||
|
token=token,
|
||||||
|
database_path=database_path,
|
||||||
|
ingest_only=ingest_only,
|
||||||
|
metrics_address=metrics_address,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def _run(config: _CliConfig) -> signal.Signals | None:
|
||||||
|
"""Run the bot under CLI-controlled process orchestration."""
|
||||||
|
bot = Crabstero(
|
||||||
|
database_path=config.database_path,
|
||||||
|
ingest_only=config.ingest_only,
|
||||||
|
metrics_address=config.metrics_address,
|
||||||
|
)
|
||||||
|
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
watchdog_stop = asyncio.Event()
|
||||||
|
watchdog_task: asyncio.Task[None] | None = None
|
||||||
|
shutdown_task: asyncio.Task[None] | None = None
|
||||||
|
ready_notified = False
|
||||||
|
stopping_notified = False
|
||||||
|
cleanup_started = False
|
||||||
|
shutdown_signal: signal.Signals | None = None
|
||||||
|
runner_task = asyncio.current_task()
|
||||||
|
if runner_task is None:
|
||||||
|
msg = "Crabstero runner is not running in a task"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
|
||||||
|
def notify_stopping() -> None:
|
||||||
|
nonlocal stopping_notified
|
||||||
|
if stopping_notified:
|
||||||
|
return
|
||||||
|
stopping_notified = True
|
||||||
|
_sd_notify("STOPPING=1")
|
||||||
|
|
||||||
|
async def shutdown(sig: signal.Signals) -> None:
|
||||||
|
logger.info("Received shutdown signal %s.", sig.name)
|
||||||
|
notify_stopping()
|
||||||
|
await bot.close()
|
||||||
|
|
||||||
|
def request_shutdown(sig: signal.Signals) -> None:
|
||||||
|
nonlocal shutdown_signal, shutdown_task
|
||||||
|
if shutdown_signal is not None:
|
||||||
|
return
|
||||||
|
|
||||||
|
shutdown_signal = sig
|
||||||
|
if not cleanup_started:
|
||||||
|
shutdown_task = asyncio.create_task(
|
||||||
|
shutdown(sig),
|
||||||
|
name="crabstero-shutdown",
|
||||||
|
)
|
||||||
|
runner_task.cancel()
|
||||||
|
|
||||||
|
async def cleanup() -> None:
|
||||||
|
nonlocal cleanup_started
|
||||||
|
if cleanup_started:
|
||||||
|
return
|
||||||
|
|
||||||
|
cleanup_started = True
|
||||||
|
|
||||||
|
if ready_notified:
|
||||||
|
notify_stopping()
|
||||||
|
|
||||||
|
watchdog_stop.set()
|
||||||
|
if watchdog_task is not None:
|
||||||
|
watchdog_task.cancel()
|
||||||
|
with contextlib.suppress(asyncio.CancelledError):
|
||||||
|
await watchdog_task
|
||||||
|
|
||||||
|
if shutdown_task is not None:
|
||||||
|
await shutdown_task
|
||||||
|
|
||||||
|
with _install_signal_handlers(loop, request_shutdown):
|
||||||
|
try:
|
||||||
|
async with bot:
|
||||||
|
try:
|
||||||
|
await bot.login(config.token)
|
||||||
|
if bot.is_closed() or stopping_notified:
|
||||||
|
return shutdown_signal
|
||||||
|
|
||||||
|
_sd_notify("READY=1")
|
||||||
|
ready_notified = True
|
||||||
|
|
||||||
|
wd_interval = _watchdog_interval()
|
||||||
|
if wd_interval is not None:
|
||||||
|
watchdog_interval = wd_interval
|
||||||
|
logger.info(
|
||||||
|
"systemd watchdog enabled (pinging every %.1fs).",
|
||||||
|
watchdog_interval,
|
||||||
|
)
|
||||||
|
watchdog_task = asyncio.create_task(
|
||||||
|
_watchdog_loop(watchdog_interval, watchdog_stop),
|
||||||
|
name="systemd-watchdog",
|
||||||
|
)
|
||||||
|
|
||||||
|
await bot.connect(reconnect=True)
|
||||||
|
finally:
|
||||||
|
await cleanup()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
if shutdown_signal is None:
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
await cleanup()
|
||||||
|
|
||||||
|
return shutdown_signal
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: list[str] | None = None) -> int:
|
||||||
|
"""Run the bot CLI.
|
||||||
|
|
||||||
|
:param argv: Optional argument list (defaults to sys.argv[1:]).
|
||||||
|
:return: Process exit code.
|
||||||
|
"""
|
||||||
|
config = _parse_args(argv)
|
||||||
|
|
||||||
|
logging.basicConfig(
|
||||||
|
level=logging.INFO,
|
||||||
|
format="[%(asctime)s] [%(name)s] [%(levelname)s] %(message)s",
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
shutdown_signal = uvloop.run(_run(config))
|
||||||
|
except KeyboardInterrupt:
|
||||||
|
logger.info("Interrupted.")
|
||||||
|
return _KEYBOARD_INTERRUPT_EXIT_CODE
|
||||||
|
|
||||||
|
if shutdown_signal == signal.SIGINT:
|
||||||
|
return _KEYBOARD_INTERRUPT_EXIT_CODE
|
||||||
|
return 0
|
||||||
@@ -0,0 +1,826 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Async SQLite database access for Crabstero's persistent storage."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from collections import Counter
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
|
from typing import TYPE_CHECKING, NamedTuple, Self, cast
|
||||||
|
|
||||||
|
import aiosqlite
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import AsyncIterator, Sequence
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class StartWord(NamedTuple):
|
||||||
|
"""A Markov chain starting word row."""
|
||||||
|
|
||||||
|
channel_id: int
|
||||||
|
user_id: int
|
||||||
|
word: str
|
||||||
|
|
||||||
|
|
||||||
|
class Transition(NamedTuple):
|
||||||
|
"""A Markov chain word transition row."""
|
||||||
|
|
||||||
|
channel_id: int
|
||||||
|
user_id: int
|
||||||
|
word: str
|
||||||
|
next_word: str
|
||||||
|
|
||||||
|
|
||||||
|
class ChannelImage(NamedTuple):
|
||||||
|
"""An image URL associated with a channel."""
|
||||||
|
|
||||||
|
channel_id: int
|
||||||
|
user_id: int
|
||||||
|
url: str
|
||||||
|
|
||||||
|
|
||||||
|
# SQLite schema version. Version 0 is the original duplicate-row Markov schema.
|
||||||
|
_CURRENT_SCHEMA_VERSION = 1
|
||||||
|
|
||||||
|
|
||||||
|
_CURRENT_MARKOV_TABLES_SQL = """
|
||||||
|
-- Markov chain starting words.
|
||||||
|
-- Occurrence counters represent frequency weight per contributing user.
|
||||||
|
CREATE TABLE IF NOT EXISTS markov_start_words (
|
||||||
|
channel_id INTEGER NOT NULL,
|
||||||
|
user_id INTEGER NOT NULL,
|
||||||
|
word TEXT NOT NULL,
|
||||||
|
occurrences INTEGER NOT NULL CHECK (occurrences > 0),
|
||||||
|
PRIMARY KEY (channel_id, word, user_id)
|
||||||
|
) WITHOUT ROWID;
|
||||||
|
|
||||||
|
-- Markov chain word transitions.
|
||||||
|
-- Occurrence counters represent frequency weight per contributing user.
|
||||||
|
CREATE TABLE IF NOT EXISTS markov_transitions (
|
||||||
|
channel_id INTEGER NOT NULL,
|
||||||
|
user_id INTEGER NOT NULL,
|
||||||
|
word TEXT NOT NULL,
|
||||||
|
next_word TEXT NOT NULL,
|
||||||
|
occurrences INTEGER NOT NULL CHECK (occurrences > 0),
|
||||||
|
PRIMARY KEY (channel_id, word, next_word, user_id)
|
||||||
|
) WITHOUT ROWID;
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
_CURRENT_NON_MARKOV_TABLES_SQL = """
|
||||||
|
|
||||||
|
-- Image URLs per channel.
|
||||||
|
CREATE TABLE IF NOT EXISTS channel_images (
|
||||||
|
channel_id INTEGER NOT NULL,
|
||||||
|
user_id INTEGER NOT NULL,
|
||||||
|
url TEXT NOT NULL
|
||||||
|
);
|
||||||
|
|
||||||
|
-- Flags for channels, servers, and users.
|
||||||
|
CREATE TABLE IF NOT EXISTS flags (
|
||||||
|
entity_type TEXT NOT NULL,
|
||||||
|
entity_id TEXT NOT NULL,
|
||||||
|
flag_name TEXT NOT NULL,
|
||||||
|
PRIMARY KEY (entity_type, entity_id, flag_name)
|
||||||
|
);
|
||||||
|
|
||||||
|
-- Tracks which channels have been bulk-ingested.
|
||||||
|
CREATE TABLE IF NOT EXISTS ingested_channels (
|
||||||
|
channel_id INTEGER NOT NULL PRIMARY KEY
|
||||||
|
);
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
_CURRENT_INDEXES_SQL = """
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_start_user ON markov_start_words(user_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_transitions_user ON markov_transitions(user_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_images_channel ON channel_images(channel_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_images_user ON channel_images(user_id);
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
_CREATE_CURRENT_SCHEMA_OBJECTS_SQL = (
|
||||||
|
"""
|
||||||
|
BEGIN IMMEDIATE;
|
||||||
|
"""
|
||||||
|
+ _CURRENT_MARKOV_TABLES_SQL
|
||||||
|
+ _CURRENT_NON_MARKOV_TABLES_SQL
|
||||||
|
+ _CURRENT_INDEXES_SQL
|
||||||
|
+ """
|
||||||
|
COMMIT;
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_MIGRATE_LEGACY_MARKOV_TO_COUNTED_ROWS_SQL = (
|
||||||
|
"""
|
||||||
|
BEGIN IMMEDIATE;
|
||||||
|
|
||||||
|
ALTER TABLE markov_start_words RENAME TO markov_start_words_v0;
|
||||||
|
ALTER TABLE markov_transitions RENAME TO markov_transitions_v0;
|
||||||
|
|
||||||
|
"""
|
||||||
|
+ _CURRENT_MARKOV_TABLES_SQL
|
||||||
|
+ _CURRENT_NON_MARKOV_TABLES_SQL
|
||||||
|
+ """
|
||||||
|
|
||||||
|
INSERT INTO markov_start_words (channel_id, user_id, word, occurrences)
|
||||||
|
SELECT channel_id, user_id, word, COUNT(*)
|
||||||
|
FROM markov_start_words_v0
|
||||||
|
GROUP BY channel_id, word, user_id
|
||||||
|
ORDER BY channel_id, word, user_id;
|
||||||
|
|
||||||
|
INSERT INTO markov_transitions (
|
||||||
|
channel_id,
|
||||||
|
user_id,
|
||||||
|
word,
|
||||||
|
next_word,
|
||||||
|
occurrences
|
||||||
|
)
|
||||||
|
SELECT channel_id, user_id, word, next_word, COUNT(*)
|
||||||
|
FROM markov_transitions_v0
|
||||||
|
GROUP BY channel_id, word, next_word, user_id
|
||||||
|
ORDER BY channel_id, word, next_word, user_id;
|
||||||
|
|
||||||
|
DROP TABLE markov_start_words_v0;
|
||||||
|
DROP TABLE markov_transitions_v0;
|
||||||
|
|
||||||
|
"""
|
||||||
|
+ _CURRENT_INDEXES_SQL
|
||||||
|
+ """
|
||||||
|
COMMIT;
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _has_expected_markov_primary_keys(
|
||||||
|
start_columns: Sequence[Sequence[object]],
|
||||||
|
transition_columns: Sequence[Sequence[object]],
|
||||||
|
) -> bool:
|
||||||
|
"""Return whether the Markov tables have the optimized primary keys."""
|
||||||
|
start_primary_keys = {str(row[1]): cast("int", row[5]) for row in start_columns}
|
||||||
|
transition_primary_keys = {
|
||||||
|
str(row[1]): cast("int", row[5]) for row in transition_columns
|
||||||
|
}
|
||||||
|
return start_primary_keys == {
|
||||||
|
"channel_id": 1,
|
||||||
|
"word": 2,
|
||||||
|
"user_id": 3,
|
||||||
|
"occurrences": 0,
|
||||||
|
} and transition_primary_keys == {
|
||||||
|
"channel_id": 1,
|
||||||
|
"word": 2,
|
||||||
|
"next_word": 3,
|
||||||
|
"user_id": 4,
|
||||||
|
"occurrences": 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _get_table_sql(connection: aiosqlite.Connection, table_name: str) -> str:
|
||||||
|
"""Return the SQLite schema SQL for a table."""
|
||||||
|
async with connection.execute(
|
||||||
|
"SELECT sql FROM sqlite_master WHERE type = 'table' AND name = ?",
|
||||||
|
(table_name,),
|
||||||
|
) as cursor:
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
return str(row[0]) if row else ""
|
||||||
|
|
||||||
|
|
||||||
|
class Database:
|
||||||
|
"""Manages all SQLite database operations for Crabstero.
|
||||||
|
|
||||||
|
Uses aiosqlite for native async access. A single connection is held
|
||||||
|
open for the lifetime of the bot process.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, connection: aiosqlite.Connection) -> None:
|
||||||
|
"""Initialize the Database wrapper with an already-opened aiosqlite connection.
|
||||||
|
|
||||||
|
:param connection: An open aiosqlite connection.
|
||||||
|
"""
|
||||||
|
self._connection = connection
|
||||||
|
self._tx_lock = asyncio.Lock()
|
||||||
|
|
||||||
|
@asynccontextmanager
|
||||||
|
async def _transaction(self) -> AsyncIterator[None]:
|
||||||
|
"""Begin an immediate write transaction.
|
||||||
|
|
||||||
|
Acquires an async lock so that only one transaction runs at a
|
||||||
|
time, then commits on success or rolls back on error.
|
||||||
|
"""
|
||||||
|
async with self._tx_lock:
|
||||||
|
await self._connection.execute("BEGIN IMMEDIATE")
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
except BaseException:
|
||||||
|
await self._connection.rollback()
|
||||||
|
raise
|
||||||
|
else:
|
||||||
|
await self._connection.commit()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def connect(cls, path: str) -> Self:
|
||||||
|
"""Open a SQLite database, configure it, and create the schema.
|
||||||
|
|
||||||
|
Configure the database for performance and create the schema if it
|
||||||
|
does not already exist.
|
||||||
|
|
||||||
|
:param path: The file path to the SQLite database.
|
||||||
|
:return: A new Database instance ready for use.
|
||||||
|
"""
|
||||||
|
connection = await aiosqlite.connect(path, isolation_level=None)
|
||||||
|
try:
|
||||||
|
# Set synchronous to NORMAL for a balance between safety and speed.
|
||||||
|
await connection.execute("PRAGMA synchronous=NORMAL")
|
||||||
|
|
||||||
|
async with connection.execute("PRAGMA user_version") as cursor:
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
version = int(row[0]) if row else 0
|
||||||
|
if version > _CURRENT_SCHEMA_VERSION:
|
||||||
|
msg = (
|
||||||
|
f"Database schema version {version} is newer than supported "
|
||||||
|
f"version {_CURRENT_SCHEMA_VERSION}."
|
||||||
|
)
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
|
||||||
|
async with connection.execute(
|
||||||
|
"SELECT name FROM sqlite_master"
|
||||||
|
" WHERE type = 'table'"
|
||||||
|
" AND name IN ('markov_start_words', 'markov_transitions')",
|
||||||
|
) as cursor:
|
||||||
|
markov_tables = {str(row[0]) for row in await cursor.fetchall()}
|
||||||
|
has_start_words = "markov_start_words" in markov_tables
|
||||||
|
has_transitions = "markov_transitions" in markov_tables
|
||||||
|
if has_start_words != has_transitions:
|
||||||
|
raise RuntimeError("Database schema is missing a Markov table.")
|
||||||
|
|
||||||
|
has_counted_markov = False
|
||||||
|
has_current_markov = False
|
||||||
|
if has_start_words:
|
||||||
|
async with connection.execute(
|
||||||
|
"PRAGMA table_info(markov_start_words)",
|
||||||
|
) as cursor:
|
||||||
|
start_column_rows = [tuple(row) for row in await cursor.fetchall()]
|
||||||
|
async with connection.execute(
|
||||||
|
"PRAGMA table_info(markov_transitions)",
|
||||||
|
) as cursor:
|
||||||
|
transition_column_rows = [
|
||||||
|
tuple(row) for row in await cursor.fetchall()
|
||||||
|
]
|
||||||
|
start_columns = {str(row[1]) for row in start_column_rows}
|
||||||
|
transition_columns = {str(row[1]) for row in transition_column_rows}
|
||||||
|
start_has_occurrences = "occurrences" in start_columns
|
||||||
|
transition_has_occurrences = "occurrences" in transition_columns
|
||||||
|
if start_has_occurrences != transition_has_occurrences:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Database schema has mismatched Markov occurrence columns.",
|
||||||
|
)
|
||||||
|
has_counted_markov = (
|
||||||
|
start_has_occurrences and transition_has_occurrences
|
||||||
|
)
|
||||||
|
if has_counted_markov:
|
||||||
|
start_table_sql = await _get_table_sql(
|
||||||
|
connection,
|
||||||
|
"markov_start_words",
|
||||||
|
)
|
||||||
|
transition_table_sql = await _get_table_sql(
|
||||||
|
connection,
|
||||||
|
"markov_transitions",
|
||||||
|
)
|
||||||
|
has_current_markov = (
|
||||||
|
_has_expected_markov_primary_keys(
|
||||||
|
start_column_rows,
|
||||||
|
transition_column_rows,
|
||||||
|
)
|
||||||
|
and "WITHOUT ROWID" in start_table_sql.upper()
|
||||||
|
and "WITHOUT ROWID" in transition_table_sql.upper()
|
||||||
|
)
|
||||||
|
if not has_current_markov:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Database schema has unsupported counted Markov tables.",
|
||||||
|
)
|
||||||
|
|
||||||
|
needs_markov_count_migration = has_start_words and not has_counted_markov
|
||||||
|
needs_post_migration_vacuum = (
|
||||||
|
has_current_markov and version < _CURRENT_SCHEMA_VERSION
|
||||||
|
)
|
||||||
|
if version == _CURRENT_SCHEMA_VERSION and needs_markov_count_migration:
|
||||||
|
raise RuntimeError(
|
||||||
|
"Database schema version 1 has legacy Markov tables.",
|
||||||
|
)
|
||||||
|
|
||||||
|
if not needs_markov_count_migration:
|
||||||
|
schema_script = _CREATE_CURRENT_SCHEMA_OBJECTS_SQL
|
||||||
|
else:
|
||||||
|
logger.info(
|
||||||
|
"Migrating legacy Markov tables to counted occurrences; "
|
||||||
|
"startup will continue after migration completes.",
|
||||||
|
)
|
||||||
|
schema_script = _MIGRATE_LEGACY_MARKOV_TO_COUNTED_ROWS_SQL
|
||||||
|
|
||||||
|
try:
|
||||||
|
await connection.executescript(schema_script)
|
||||||
|
except BaseException:
|
||||||
|
await connection.rollback()
|
||||||
|
raise
|
||||||
|
if needs_markov_count_migration:
|
||||||
|
logger.info("Legacy Markov table migration complete.")
|
||||||
|
if needs_markov_count_migration or needs_post_migration_vacuum:
|
||||||
|
if needs_post_migration_vacuum:
|
||||||
|
logger.info(
|
||||||
|
"Retrying post-migration database vacuum before marking "
|
||||||
|
"schema version current.",
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"Vacuuming database after legacy Markov table migration; "
|
||||||
|
"startup will continue after vacuum completes.",
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
async with connection.execute("VACUUM"):
|
||||||
|
pass
|
||||||
|
except Exception:
|
||||||
|
logger.exception(
|
||||||
|
"Database vacuum failed after legacy Markov table "
|
||||||
|
"migration; continuing with migrated database.",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.info("Post-migration database vacuum complete.")
|
||||||
|
await connection.execute(
|
||||||
|
f"PRAGMA user_version = {_CURRENT_SCHEMA_VERSION}",
|
||||||
|
)
|
||||||
|
elif version < _CURRENT_SCHEMA_VERSION:
|
||||||
|
await connection.execute(
|
||||||
|
f"PRAGMA user_version = {_CURRENT_SCHEMA_VERSION}",
|
||||||
|
)
|
||||||
|
except BaseException:
|
||||||
|
await connection.close()
|
||||||
|
raise
|
||||||
|
|
||||||
|
return cls(connection)
|
||||||
|
|
||||||
|
async def close(self) -> None:
|
||||||
|
"""Close the database connection."""
|
||||||
|
await self._connection.close()
|
||||||
|
|
||||||
|
async def add_markov_data(
|
||||||
|
self,
|
||||||
|
start_words: list[StartWord],
|
||||||
|
transitions: list[Transition],
|
||||||
|
) -> None:
|
||||||
|
"""Insert or increment Markov start words and transitions, then commit.
|
||||||
|
|
||||||
|
Both writes happen in a single transaction.
|
||||||
|
|
||||||
|
:param start_words: Starting word entries to insert.
|
||||||
|
:param transitions: Transition entries to insert.
|
||||||
|
"""
|
||||||
|
if start_words or transitions:
|
||||||
|
async with self._transaction():
|
||||||
|
if start_words:
|
||||||
|
counted_start_words = [
|
||||||
|
(*start_word, occurrences)
|
||||||
|
for start_word, occurrences in Counter(start_words).items()
|
||||||
|
]
|
||||||
|
await self._connection.executemany(
|
||||||
|
"INSERT INTO markov_start_words"
|
||||||
|
" (channel_id, user_id, word, occurrences)"
|
||||||
|
" VALUES (?, ?, ?, ?)"
|
||||||
|
" ON CONFLICT(channel_id, word, user_id)"
|
||||||
|
" DO UPDATE SET occurrences ="
|
||||||
|
" markov_start_words.occurrences + excluded.occurrences",
|
||||||
|
counted_start_words,
|
||||||
|
)
|
||||||
|
if transitions:
|
||||||
|
counted_transitions = [
|
||||||
|
(*transition, occurrences)
|
||||||
|
for transition, occurrences in Counter(transitions).items()
|
||||||
|
]
|
||||||
|
await self._connection.executemany(
|
||||||
|
"INSERT INTO markov_transitions"
|
||||||
|
" (channel_id, user_id, word, next_word, occurrences)"
|
||||||
|
" VALUES (?, ?, ?, ?, ?)"
|
||||||
|
" ON CONFLICT(channel_id, word, next_word, user_id)"
|
||||||
|
" DO UPDATE SET occurrences ="
|
||||||
|
" markov_transitions.occurrences + excluded.occurrences",
|
||||||
|
counted_transitions,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def remove_markov_data(
|
||||||
|
self,
|
||||||
|
start_words: list[StartWord],
|
||||||
|
transitions: list[Transition],
|
||||||
|
) -> None:
|
||||||
|
"""Remove one matching occurrence per entry from the Markov tables.
|
||||||
|
|
||||||
|
Each entry removes at most one counted occurrence, preserving remaining
|
||||||
|
frequency weight.
|
||||||
|
|
||||||
|
:param start_words: Starting word entries to remove.
|
||||||
|
:param transitions: Transition entries to remove.
|
||||||
|
"""
|
||||||
|
if start_words or transitions:
|
||||||
|
async with self._transaction():
|
||||||
|
if start_words:
|
||||||
|
counted_start_words = [
|
||||||
|
(*start_word, occurrences)
|
||||||
|
for start_word, occurrences in Counter(start_words).items()
|
||||||
|
]
|
||||||
|
await self._connection.executemany(
|
||||||
|
"DELETE FROM markov_start_words"
|
||||||
|
" WHERE channel_id = ?"
|
||||||
|
" AND user_id = ?"
|
||||||
|
" AND word = ?"
|
||||||
|
" AND occurrences <= ?",
|
||||||
|
counted_start_words,
|
||||||
|
)
|
||||||
|
await self._connection.executemany(
|
||||||
|
"UPDATE markov_start_words"
|
||||||
|
" SET occurrences = occurrences - ?"
|
||||||
|
" WHERE channel_id = ?"
|
||||||
|
" AND user_id = ?"
|
||||||
|
" AND word = ?"
|
||||||
|
" AND occurrences > ?",
|
||||||
|
[
|
||||||
|
(
|
||||||
|
occurrences,
|
||||||
|
start_word.channel_id,
|
||||||
|
start_word.user_id,
|
||||||
|
start_word.word,
|
||||||
|
occurrences,
|
||||||
|
)
|
||||||
|
for start_word, occurrences in Counter(start_words).items()
|
||||||
|
],
|
||||||
|
)
|
||||||
|
if transitions:
|
||||||
|
counted_transitions = [
|
||||||
|
(*transition, occurrences)
|
||||||
|
for transition, occurrences in Counter(transitions).items()
|
||||||
|
]
|
||||||
|
await self._connection.executemany(
|
||||||
|
"DELETE FROM markov_transitions"
|
||||||
|
" WHERE channel_id = ?"
|
||||||
|
" AND user_id = ?"
|
||||||
|
" AND word = ?"
|
||||||
|
" AND next_word = ?"
|
||||||
|
" AND occurrences <= ?",
|
||||||
|
counted_transitions,
|
||||||
|
)
|
||||||
|
await self._connection.executemany(
|
||||||
|
"UPDATE markov_transitions"
|
||||||
|
" SET occurrences = occurrences - ?"
|
||||||
|
" WHERE channel_id = ?"
|
||||||
|
" AND user_id = ?"
|
||||||
|
" AND word = ?"
|
||||||
|
" AND next_word = ?"
|
||||||
|
" AND occurrences > ?",
|
||||||
|
[
|
||||||
|
(
|
||||||
|
occurrences,
|
||||||
|
transition.channel_id,
|
||||||
|
transition.user_id,
|
||||||
|
transition.word,
|
||||||
|
transition.next_word,
|
||||||
|
occurrences,
|
||||||
|
)
|
||||||
|
for transition, occurrences in Counter(transitions).items()
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
async def get_random_start_word(self, channel_id: int) -> str | None:
|
||||||
|
"""Return a random starting word for a channel.
|
||||||
|
|
||||||
|
Weighted by occurrence frequency.
|
||||||
|
|
||||||
|
:param channel_id: The Discord channel ID.
|
||||||
|
:return: A random starting word, or None if none exist.
|
||||||
|
"""
|
||||||
|
# Mask RANDOM() instead of ABS(RANDOM()) to avoid signed 64-bit overflow.
|
||||||
|
async with self._connection.execute(
|
||||||
|
"""
|
||||||
|
WITH weighted AS (
|
||||||
|
SELECT word, SUM(occurrences) AS weight
|
||||||
|
FROM markov_start_words
|
||||||
|
WHERE channel_id = ?
|
||||||
|
GROUP BY word
|
||||||
|
),
|
||||||
|
total AS (
|
||||||
|
SELECT SUM(weight) AS total_weight
|
||||||
|
FROM weighted
|
||||||
|
),
|
||||||
|
choice AS MATERIALIZED (
|
||||||
|
SELECT
|
||||||
|
((RANDOM() & 9223372036854775807) % total_weight) + 1
|
||||||
|
AS selected_weight
|
||||||
|
FROM total
|
||||||
|
WHERE total_weight > 0
|
||||||
|
),
|
||||||
|
ranked AS (
|
||||||
|
SELECT
|
||||||
|
word,
|
||||||
|
SUM(weight) OVER (
|
||||||
|
ORDER BY word
|
||||||
|
ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW
|
||||||
|
) AS cumulative_weight
|
||||||
|
FROM weighted
|
||||||
|
)
|
||||||
|
SELECT word
|
||||||
|
FROM ranked
|
||||||
|
CROSS JOIN choice
|
||||||
|
WHERE cumulative_weight >= selected_weight
|
||||||
|
ORDER BY cumulative_weight
|
||||||
|
LIMIT 1
|
||||||
|
""",
|
||||||
|
(channel_id,),
|
||||||
|
) as cursor:
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
return row[0] if row else None
|
||||||
|
|
||||||
|
async def get_random_next_word(self, channel_id: int, word: str) -> str | None:
|
||||||
|
"""Return a random next word for a given word in a channel.
|
||||||
|
|
||||||
|
Weighted by occurrence frequency.
|
||||||
|
|
||||||
|
:param channel_id: The Discord channel ID.
|
||||||
|
:param word: The current word to find a transition for.
|
||||||
|
:return: A random next word, or None if no transitions exist.
|
||||||
|
"""
|
||||||
|
# Mask RANDOM() instead of ABS(RANDOM()) to avoid signed 64-bit overflow.
|
||||||
|
async with self._connection.execute(
|
||||||
|
"""
|
||||||
|
WITH weighted AS (
|
||||||
|
SELECT next_word, SUM(occurrences) AS weight
|
||||||
|
FROM markov_transitions
|
||||||
|
WHERE channel_id = ? AND word = ?
|
||||||
|
GROUP BY next_word
|
||||||
|
),
|
||||||
|
total AS (
|
||||||
|
SELECT SUM(weight) AS total_weight
|
||||||
|
FROM weighted
|
||||||
|
),
|
||||||
|
choice AS MATERIALIZED (
|
||||||
|
SELECT
|
||||||
|
((RANDOM() & 9223372036854775807) % total_weight) + 1
|
||||||
|
AS selected_weight
|
||||||
|
FROM total
|
||||||
|
WHERE total_weight > 0
|
||||||
|
),
|
||||||
|
ranked AS (
|
||||||
|
SELECT
|
||||||
|
next_word,
|
||||||
|
SUM(weight) OVER (
|
||||||
|
ORDER BY next_word
|
||||||
|
ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW
|
||||||
|
) AS cumulative_weight
|
||||||
|
FROM weighted
|
||||||
|
)
|
||||||
|
SELECT next_word
|
||||||
|
FROM ranked
|
||||||
|
CROSS JOIN choice
|
||||||
|
WHERE cumulative_weight >= selected_weight
|
||||||
|
ORDER BY cumulative_weight
|
||||||
|
LIMIT 1
|
||||||
|
""",
|
||||||
|
(channel_id, word),
|
||||||
|
) as cursor:
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
return row[0] if row else None
|
||||||
|
|
||||||
|
async def get_random_completing_next_word(
|
||||||
|
self,
|
||||||
|
channel_id: int,
|
||||||
|
word: str,
|
||||||
|
) -> str | None:
|
||||||
|
"""Return a random sentence-ending next word for a given word in a channel.
|
||||||
|
|
||||||
|
Weighted by occurrence frequency. A completing word is one whose last
|
||||||
|
character is '.', '!', '?', or '§'.
|
||||||
|
|
||||||
|
The set of terminators must match ``markov._TERMINATORS``.
|
||||||
|
|
||||||
|
:param channel_id: The Discord channel ID.
|
||||||
|
:param word: The current word to find a completing transition for.
|
||||||
|
:return: A random completing next word, or None if none exist.
|
||||||
|
"""
|
||||||
|
# Mask RANDOM() instead of ABS(RANDOM()) to avoid signed 64-bit overflow.
|
||||||
|
async with self._connection.execute(
|
||||||
|
"""
|
||||||
|
WITH weighted AS (
|
||||||
|
SELECT next_word, SUM(occurrences) AS weight
|
||||||
|
FROM markov_transitions
|
||||||
|
WHERE channel_id = ?
|
||||||
|
AND word = ?
|
||||||
|
AND SUBSTR(next_word, -1, 1) IN ('.', '!', '?', '§')
|
||||||
|
GROUP BY next_word
|
||||||
|
),
|
||||||
|
total AS (
|
||||||
|
SELECT SUM(weight) AS total_weight
|
||||||
|
FROM weighted
|
||||||
|
),
|
||||||
|
choice AS MATERIALIZED (
|
||||||
|
SELECT
|
||||||
|
((RANDOM() & 9223372036854775807) % total_weight) + 1
|
||||||
|
AS selected_weight
|
||||||
|
FROM total
|
||||||
|
WHERE total_weight > 0
|
||||||
|
),
|
||||||
|
ranked AS (
|
||||||
|
SELECT
|
||||||
|
next_word,
|
||||||
|
SUM(weight) OVER (
|
||||||
|
ORDER BY next_word
|
||||||
|
ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW
|
||||||
|
) AS cumulative_weight
|
||||||
|
FROM weighted
|
||||||
|
)
|
||||||
|
SELECT next_word
|
||||||
|
FROM ranked
|
||||||
|
CROSS JOIN choice
|
||||||
|
WHERE cumulative_weight >= selected_weight
|
||||||
|
ORDER BY cumulative_weight
|
||||||
|
LIMIT 1
|
||||||
|
""",
|
||||||
|
(channel_id, word),
|
||||||
|
) as cursor:
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
return row[0] if row else None
|
||||||
|
|
||||||
|
async def add_images(self, images: list[ChannelImage]) -> None:
|
||||||
|
"""Store image URLs for a given channel.
|
||||||
|
|
||||||
|
:param images: Image entries to insert.
|
||||||
|
"""
|
||||||
|
if images:
|
||||||
|
async with self._transaction():
|
||||||
|
await self._connection.executemany(
|
||||||
|
"INSERT INTO channel_images"
|
||||||
|
" (channel_id, user_id, url)"
|
||||||
|
" VALUES (?, ?, ?)",
|
||||||
|
images,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def remove_images(self, images: list[ChannelImage]) -> None:
|
||||||
|
"""Remove one matching row per entry from the images table.
|
||||||
|
|
||||||
|
Each entry removes at most one duplicate row, preserving remaining
|
||||||
|
frequency weight.
|
||||||
|
|
||||||
|
:param images: Image entries to remove.
|
||||||
|
"""
|
||||||
|
if images:
|
||||||
|
async with self._transaction():
|
||||||
|
await self._connection.executemany(
|
||||||
|
"DELETE FROM channel_images"
|
||||||
|
" WHERE rowid = ("
|
||||||
|
" SELECT rowid FROM channel_images"
|
||||||
|
" WHERE channel_id = ?"
|
||||||
|
" AND user_id = ?"
|
||||||
|
" AND url = ?"
|
||||||
|
" LIMIT 1"
|
||||||
|
" )",
|
||||||
|
images,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def get_random_image(self, channel_id: int) -> str | None:
|
||||||
|
"""Return a random image URL for a given channel.
|
||||||
|
|
||||||
|
:param channel_id: The Discord channel ID.
|
||||||
|
:return: A random image URL, or None if none exist.
|
||||||
|
"""
|
||||||
|
async with self._connection.execute(
|
||||||
|
"SELECT url FROM channel_images"
|
||||||
|
" WHERE channel_id = ?"
|
||||||
|
" ORDER BY RANDOM() LIMIT 1",
|
||||||
|
(channel_id,),
|
||||||
|
) as cursor:
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
return row[0] if row else None
|
||||||
|
|
||||||
|
async def set_flag(self, entity_type: str, entity_id: str, flag_name: str) -> None:
|
||||||
|
"""Set a flag on a given entity. If the flag is already set, this is a no-op.
|
||||||
|
|
||||||
|
:param entity_type: The entity type ("channel", "server", "user").
|
||||||
|
:param entity_id: The Discord ID of the entity.
|
||||||
|
:param flag_name: The name of the flag to set.
|
||||||
|
"""
|
||||||
|
async with self._transaction():
|
||||||
|
await self._connection.execute(
|
||||||
|
"INSERT OR IGNORE INTO flags"
|
||||||
|
" (entity_type, entity_id, flag_name)"
|
||||||
|
" VALUES (?, ?, ?)",
|
||||||
|
(entity_type, entity_id, flag_name),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def clear_flag(
|
||||||
|
self,
|
||||||
|
entity_type: str,
|
||||||
|
entity_id: str,
|
||||||
|
flag_name: str,
|
||||||
|
) -> None:
|
||||||
|
"""Clear a flag on a given entity. If the flag is not set, this is a no-op.
|
||||||
|
|
||||||
|
:param entity_type: The entity type ("channel", "server", "user").
|
||||||
|
:param entity_id: The Discord ID of the entity.
|
||||||
|
:param flag_name: The name of the flag to clear.
|
||||||
|
"""
|
||||||
|
async with self._transaction():
|
||||||
|
await self._connection.execute(
|
||||||
|
"DELETE FROM flags"
|
||||||
|
" WHERE entity_type = ?"
|
||||||
|
" AND entity_id = ?"
|
||||||
|
" AND flag_name = ?",
|
||||||
|
(entity_type, entity_id, flag_name),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def is_flag_set(
|
||||||
|
self,
|
||||||
|
entity_type: str,
|
||||||
|
entity_id: str,
|
||||||
|
flag_name: str,
|
||||||
|
) -> bool:
|
||||||
|
"""Check whether a flag is set on a given entity.
|
||||||
|
|
||||||
|
:param entity_type: The entity type ("channel", "server", "user").
|
||||||
|
:param entity_id: The Discord ID of the entity.
|
||||||
|
:param flag_name: The name of the flag to check.
|
||||||
|
:return: True if the flag is set, False otherwise.
|
||||||
|
"""
|
||||||
|
async with self._connection.execute(
|
||||||
|
"SELECT 1 FROM flags"
|
||||||
|
" WHERE entity_type = ?"
|
||||||
|
" AND entity_id = ?"
|
||||||
|
" AND flag_name = ?",
|
||||||
|
(entity_type, entity_id, flag_name),
|
||||||
|
) as cursor:
|
||||||
|
return await cursor.fetchone() is not None
|
||||||
|
|
||||||
|
async def is_channel_ingested(self, channel_id: int) -> bool:
|
||||||
|
"""Check whether a channel has already been bulk-ingested.
|
||||||
|
|
||||||
|
:param channel_id: The Discord channel ID.
|
||||||
|
:return: True if the channel has been ingested, False otherwise.
|
||||||
|
"""
|
||||||
|
async with self._connection.execute(
|
||||||
|
"SELECT 1 FROM ingested_channels WHERE channel_id = ?",
|
||||||
|
(channel_id,),
|
||||||
|
) as cursor:
|
||||||
|
return await cursor.fetchone() is not None
|
||||||
|
|
||||||
|
async def mark_channel_ingested(self, channel_id: int) -> None:
|
||||||
|
"""Mark a channel as having been bulk-ingested.
|
||||||
|
|
||||||
|
:param channel_id: The Discord channel ID.
|
||||||
|
"""
|
||||||
|
async with self._transaction():
|
||||||
|
await self._connection.execute(
|
||||||
|
"INSERT OR IGNORE INTO ingested_channels (channel_id) VALUES (?)",
|
||||||
|
(channel_id,),
|
||||||
|
)
|
||||||
|
|
||||||
|
async def forget_user(self, user_id: int, no_ingest_flag: str) -> None:
|
||||||
|
"""Delete all user data, clear flags, and set noIngest atomically.
|
||||||
|
|
||||||
|
Sets the noIngest flag, clears all other user flags, and deletes
|
||||||
|
all Markov and image data in a single transaction.
|
||||||
|
|
||||||
|
:param user_id: The Discord user ID to forget.
|
||||||
|
:param no_ingest_flag: The flag name for noIngest.
|
||||||
|
"""
|
||||||
|
entity_id = str(user_id)
|
||||||
|
async with self._transaction():
|
||||||
|
await self._connection.execute(
|
||||||
|
"INSERT OR IGNORE INTO flags (entity_type, entity_id, flag_name)"
|
||||||
|
" VALUES ('user', ?, ?)",
|
||||||
|
(entity_id, no_ingest_flag),
|
||||||
|
)
|
||||||
|
await self._connection.execute(
|
||||||
|
"DELETE FROM flags WHERE entity_type = 'user'"
|
||||||
|
" AND entity_id = ? AND flag_name != ?",
|
||||||
|
(entity_id, no_ingest_flag),
|
||||||
|
)
|
||||||
|
await self._connection.execute(
|
||||||
|
"DELETE FROM markov_start_words WHERE user_id = ?",
|
||||||
|
(user_id,),
|
||||||
|
)
|
||||||
|
await self._connection.execute(
|
||||||
|
"DELETE FROM markov_transitions WHERE user_id = ?",
|
||||||
|
(user_id,),
|
||||||
|
)
|
||||||
|
await self._connection.execute(
|
||||||
|
"DELETE FROM channel_images WHERE user_id = ?",
|
||||||
|
(user_id,),
|
||||||
|
)
|
||||||
@@ -12,8 +12,7 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
"""
|
"""Provides async convenience wrappers around the database flag methods.
|
||||||
Provides async convenience wrappers around the database flag methods.
|
|
||||||
|
|
||||||
Accepts discord.py objects or raw integer IDs and translates them into the
|
Accepts discord.py objects or raw integer IDs and translates them into the
|
||||||
entity_type/entity_id pairs used by the database layer.
|
entity_type/entity_id pairs used by the database layer.
|
||||||
@@ -28,7 +27,7 @@ if TYPE_CHECKING:
|
|||||||
from crabstero.database import Database
|
from crabstero.database import Database
|
||||||
|
|
||||||
|
|
||||||
class Flag(enum.Enum):
|
class Flag(enum.StrEnum):
|
||||||
"""Enum of all supported flag names."""
|
"""Enum of all supported flag names."""
|
||||||
|
|
||||||
NO_REPLY = "noReply"
|
NO_REPLY = "noReply"
|
||||||
@@ -36,7 +35,7 @@ class Flag(enum.Enum):
|
|||||||
ALLOW_PINGS = "allowPings"
|
ALLOW_PINGS = "allowPings"
|
||||||
|
|
||||||
|
|
||||||
class EntityType(enum.Enum):
|
class EntityType(enum.StrEnum):
|
||||||
"""Enum of entity types that can have flags."""
|
"""Enum of entity types that can have flags."""
|
||||||
|
|
||||||
CHANNEL = "channel"
|
CHANNEL = "channel"
|
||||||
@@ -45,8 +44,7 @@ class EntityType(enum.Enum):
|
|||||||
|
|
||||||
|
|
||||||
def _entity_id(entity: discord.abc.Snowflake | int) -> str:
|
def _entity_id(entity: discord.abc.Snowflake | int) -> str:
|
||||||
"""
|
"""Extract a string entity ID from a Discord object or raw integer ID.
|
||||||
Extracts a string entity ID from a Discord object or raw integer ID.
|
|
||||||
|
|
||||||
:param entity: A Discord entity or raw integer ID.
|
:param entity: A Discord entity or raw integer ID.
|
||||||
:return: The entity ID as a string.
|
:return: The entity ID as a string.
|
||||||
@@ -60,15 +58,14 @@ async def set_flag(
|
|||||||
entity_type: EntityType,
|
entity_type: EntityType,
|
||||||
flag: Flag,
|
flag: Flag,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""Set a flag on a given entity.
|
||||||
Sets a flag on a given entity.
|
|
||||||
|
|
||||||
:param db: The database instance.
|
:param db: The database instance.
|
||||||
:param entity: The Discord entity or raw integer ID to set the flag on.
|
:param entity: The Discord entity or raw integer ID to set the flag on.
|
||||||
:param entity_type: The type of the entity.
|
:param entity_type: The type of the entity.
|
||||||
:param flag: The flag to set.
|
:param flag: The flag to set.
|
||||||
"""
|
"""
|
||||||
await db.set_flag(entity_type.value, _entity_id(entity), flag.value)
|
await db.set_flag(entity_type, _entity_id(entity), flag)
|
||||||
|
|
||||||
|
|
||||||
async def clear_flag(
|
async def clear_flag(
|
||||||
@@ -77,15 +74,14 @@ async def clear_flag(
|
|||||||
entity_type: EntityType,
|
entity_type: EntityType,
|
||||||
flag: Flag,
|
flag: Flag,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""Clear a flag on a given entity.
|
||||||
Clears a flag on a given entity.
|
|
||||||
|
|
||||||
:param db: The database instance.
|
:param db: The database instance.
|
||||||
:param entity: The Discord entity or raw integer ID to clear the flag on.
|
:param entity: The Discord entity or raw integer ID to clear the flag on.
|
||||||
:param entity_type: The type of the entity.
|
:param entity_type: The type of the entity.
|
||||||
:param flag: The flag to clear.
|
:param flag: The flag to clear.
|
||||||
"""
|
"""
|
||||||
await db.clear_flag(entity_type.value, _entity_id(entity), flag.value)
|
await db.clear_flag(entity_type, _entity_id(entity), flag)
|
||||||
|
|
||||||
|
|
||||||
async def is_flag_set(
|
async def is_flag_set(
|
||||||
@@ -94,8 +90,7 @@ async def is_flag_set(
|
|||||||
entity_type: EntityType,
|
entity_type: EntityType,
|
||||||
flag: Flag,
|
flag: Flag,
|
||||||
) -> bool:
|
) -> bool:
|
||||||
"""
|
"""Check whether a flag is set on a given entity.
|
||||||
Checks whether a flag is set on a given entity.
|
|
||||||
|
|
||||||
:param db: The database instance.
|
:param db: The database instance.
|
||||||
:param entity: The Discord entity or raw integer ID to check.
|
:param entity: The Discord entity or raw integer ID to check.
|
||||||
@@ -103,4 +98,4 @@ async def is_flag_set(
|
|||||||
:param flag: The flag to check for.
|
:param flag: The flag to check for.
|
||||||
:return: True if the flag is set, False otherwise.
|
:return: True if the flag is set, False otherwise.
|
||||||
"""
|
"""
|
||||||
return await db.is_flag_set(entity_type.value, _entity_id(entity), flag.value)
|
return await db.is_flag_set(entity_type, _entity_id(entity), flag)
|
||||||
@@ -12,9 +12,4 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
import crabstero
|
"""Listener cogs for Discord events."""
|
||||||
|
|
||||||
|
|
||||||
def test_version_is_available() -> None:
|
|
||||||
assert isinstance(crabstero.__version__, str)
|
|
||||||
assert crabstero.__version__
|
|
||||||
@@ -0,0 +1,225 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Handles responding to interactions.
|
||||||
|
|
||||||
|
Provides the /pingme and /forgetme slash commands as a Cog with app commands.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import TYPE_CHECKING, override
|
||||||
|
|
||||||
|
import discord
|
||||||
|
from discord import app_commands
|
||||||
|
from discord.ext import commands
|
||||||
|
|
||||||
|
from crabstero import flags, metrics
|
||||||
|
from crabstero.bot import TrackedView
|
||||||
|
from crabstero.flags import EntityType, Flag
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from crabstero.bot import Crabstero
|
||||||
|
from crabstero.database import Database
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class ForgetMeView(TrackedView):
|
||||||
|
"""Confirmation view with Confirm/Cancel buttons for /forgetme."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
user_id: int,
|
||||||
|
db: Database,
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
) -> None:
|
||||||
|
"""Create a new ForgetMeView.
|
||||||
|
|
||||||
|
:param user_id: The Discord user ID who invoked the command.
|
||||||
|
:param db: The database instance.
|
||||||
|
:param interaction: The original interaction for editing on timeout.
|
||||||
|
"""
|
||||||
|
super().__init__()
|
||||||
|
self._user_id = user_id
|
||||||
|
self._db = db
|
||||||
|
self._responded = False
|
||||||
|
self._interaction = interaction
|
||||||
|
|
||||||
|
@override
|
||||||
|
async def interaction_check(self, interaction: discord.Interaction) -> bool:
|
||||||
|
"""Only allow the invoking user to interact with the buttons.
|
||||||
|
|
||||||
|
:param interaction: The interaction event.
|
||||||
|
:return: True if the user matches, False otherwise.
|
||||||
|
"""
|
||||||
|
return interaction.user.id == self._user_id
|
||||||
|
|
||||||
|
@discord.ui.button(label="Confirm", style=discord.ButtonStyle.danger)
|
||||||
|
async def confirm(
|
||||||
|
self,
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
_button: discord.ui.Button[ForgetMeView],
|
||||||
|
) -> None:
|
||||||
|
"""Delete all user data, clear flags, and set noIngest.
|
||||||
|
|
||||||
|
:param interaction: The interaction event.
|
||||||
|
:param _button: The button that was pressed.
|
||||||
|
"""
|
||||||
|
self._responded = True
|
||||||
|
await self._db.forget_user(self._user_id, Flag.NO_INGEST)
|
||||||
|
metrics.FORGETME.labels(outcome="completed").inc()
|
||||||
|
await interaction.response.edit_message(
|
||||||
|
content=(
|
||||||
|
"I have deleted your data and I will not use"
|
||||||
|
" your messages going forward."
|
||||||
|
),
|
||||||
|
view=None,
|
||||||
|
)
|
||||||
|
self.stop()
|
||||||
|
|
||||||
|
@discord.ui.button(label="Cancel", style=discord.ButtonStyle.secondary)
|
||||||
|
async def cancel(
|
||||||
|
self,
|
||||||
|
interaction: discord.Interaction,
|
||||||
|
_button: discord.ui.Button[ForgetMeView],
|
||||||
|
) -> None:
|
||||||
|
"""Cancel the /forgetme action.
|
||||||
|
|
||||||
|
:param interaction: The interaction event.
|
||||||
|
:param _button: The button that was pressed.
|
||||||
|
"""
|
||||||
|
self._responded = True
|
||||||
|
metrics.FORGETME.labels(outcome="cancelled").inc()
|
||||||
|
await interaction.response.edit_message(
|
||||||
|
content="Action cancelled. I have not modified your data.",
|
||||||
|
view=None,
|
||||||
|
)
|
||||||
|
self.stop()
|
||||||
|
|
||||||
|
@override
|
||||||
|
async def on_timeout(self) -> None:
|
||||||
|
"""Remove buttons and increment timeout metric."""
|
||||||
|
if self._responded:
|
||||||
|
return
|
||||||
|
metrics.FORGETME.labels(outcome="timeout").inc()
|
||||||
|
try:
|
||||||
|
await self._interaction.edit_original_response(
|
||||||
|
content="This timed out. Run `/forgetme` again if you still want to.",
|
||||||
|
view=None,
|
||||||
|
)
|
||||||
|
except discord.NotFound:
|
||||||
|
logger.debug(
|
||||||
|
"Original /forgetme response for user %s was already deleted.",
|
||||||
|
self._user_id,
|
||||||
|
)
|
||||||
|
except discord.HTTPException:
|
||||||
|
logger.debug(
|
||||||
|
"Failed to edit timed-out /forgetme response for user %s.",
|
||||||
|
self._user_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class InteractionCog(commands.Cog):
|
||||||
|
"""Cog for handling slash command interactions."""
|
||||||
|
|
||||||
|
def __init__(self, bot: Crabstero) -> None:
|
||||||
|
"""Create a new interaction handler cog.
|
||||||
|
|
||||||
|
:param bot: The bot instance.
|
||||||
|
"""
|
||||||
|
self.bot = bot
|
||||||
|
|
||||||
|
@app_commands.command(
|
||||||
|
name="pingme",
|
||||||
|
description=(
|
||||||
|
"Toggle whether you receive pings for generated messages that mention you."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
async def pingme(self, interaction: discord.Interaction) -> None:
|
||||||
|
"""Toggle the allowPings flag for the user who ran the command.
|
||||||
|
|
||||||
|
:param interaction: The interaction event.
|
||||||
|
"""
|
||||||
|
if await flags.is_flag_set(
|
||||||
|
self.bot.db,
|
||||||
|
interaction.user,
|
||||||
|
EntityType.USER,
|
||||||
|
Flag.ALLOW_PINGS,
|
||||||
|
):
|
||||||
|
await flags.clear_flag(
|
||||||
|
self.bot.db,
|
||||||
|
interaction.user,
|
||||||
|
EntityType.USER,
|
||||||
|
Flag.ALLOW_PINGS,
|
||||||
|
)
|
||||||
|
metrics.PINGME.labels(outcome="opted_out").inc()
|
||||||
|
await interaction.response.send_message(
|
||||||
|
"I will no longer ping you for messages which"
|
||||||
|
" mention you. If you decide to opt back in,"
|
||||||
|
" run `/pingme` any time.",
|
||||||
|
ephemeral=True,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
await flags.set_flag(
|
||||||
|
self.bot.db,
|
||||||
|
interaction.user,
|
||||||
|
EntityType.USER,
|
||||||
|
Flag.ALLOW_PINGS,
|
||||||
|
)
|
||||||
|
metrics.PINGME.labels(outcome="opted_in").inc()
|
||||||
|
await interaction.response.send_message(
|
||||||
|
"I will now ping you for messages which mention"
|
||||||
|
" you. If you change your mind, run"
|
||||||
|
" `/pingme` any time.",
|
||||||
|
ephemeral=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
@app_commands.command(
|
||||||
|
name="forgetme",
|
||||||
|
description="Delete all your data and stop your messages from being used.",
|
||||||
|
)
|
||||||
|
async def forgetme(self, interaction: discord.Interaction) -> None:
|
||||||
|
"""Delete all user data and set the noIngest flag.
|
||||||
|
|
||||||
|
:param interaction: The interaction event.
|
||||||
|
"""
|
||||||
|
if await flags.is_flag_set(
|
||||||
|
self.bot.db,
|
||||||
|
interaction.user,
|
||||||
|
EntityType.USER,
|
||||||
|
Flag.NO_INGEST,
|
||||||
|
):
|
||||||
|
metrics.FORGETME.labels(outcome="already_forgotten").inc()
|
||||||
|
await interaction.response.send_message(
|
||||||
|
"I have already removed your data and I am not using your messages.",
|
||||||
|
ephemeral=True,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
view = ForgetMeView(interaction.user.id, self.bot.db, interaction)
|
||||||
|
await interaction.response.send_message(
|
||||||
|
"I will delete everything I have learned from your"
|
||||||
|
" messages and I will not use them going forward."
|
||||||
|
" Would you like to proceed?",
|
||||||
|
view=view,
|
||||||
|
ephemeral=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def setup(bot: Crabstero) -> None:
|
||||||
|
"""Add the InteractionCog to the bot.
|
||||||
|
|
||||||
|
:param bot: The bot instance.
|
||||||
|
"""
|
||||||
|
await bot.add_cog(InteractionCog(bot))
|
||||||
@@ -0,0 +1,111 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Handles message creation and deletion events.
|
||||||
|
|
||||||
|
Responds to mentions, ingests normal text messages, and reverses
|
||||||
|
ingestion when messages are deleted.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import discord
|
||||||
|
from discord.ext import commands
|
||||||
|
|
||||||
|
from crabstero import metrics
|
||||||
|
from crabstero.messages import ingest_message, reply_to_message, uningest_message
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from crabstero.bot import Crabstero
|
||||||
|
|
||||||
|
|
||||||
|
class MessageCog(commands.Cog):
|
||||||
|
"""Cog for handling message creation events."""
|
||||||
|
|
||||||
|
def __init__(self, bot: Crabstero) -> None:
|
||||||
|
"""Create a new message creation handler cog.
|
||||||
|
|
||||||
|
:param bot: The bot instance.
|
||||||
|
"""
|
||||||
|
self.bot = bot
|
||||||
|
|
||||||
|
@commands.Cog.listener()
|
||||||
|
async def on_message(self, message: discord.Message) -> None:
|
||||||
|
"""Respond to mentions and ingest normal text messages.
|
||||||
|
|
||||||
|
:param message: The message event.
|
||||||
|
"""
|
||||||
|
channel = message.channel
|
||||||
|
|
||||||
|
if not isinstance(
|
||||||
|
channel,
|
||||||
|
(discord.TextChannel, discord.Thread, discord.VoiceChannel),
|
||||||
|
):
|
||||||
|
return
|
||||||
|
|
||||||
|
if message.author.bot or message.author == self.bot.user:
|
||||||
|
return
|
||||||
|
|
||||||
|
metrics.MESSAGES_PROCESSED.inc()
|
||||||
|
|
||||||
|
if self.bot.ingest_only:
|
||||||
|
# In ingest-only mode, never reply — only ingest eligible messages.
|
||||||
|
if message.type == discord.MessageType.default and not isinstance(
|
||||||
|
channel,
|
||||||
|
discord.Thread,
|
||||||
|
):
|
||||||
|
await ingest_message(self.bot.db, message, self.bot.ingest_cache)
|
||||||
|
return
|
||||||
|
|
||||||
|
if self.bot.user in message.mentions:
|
||||||
|
await reply_to_message(self.bot.db, message)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Only ingest non-thread messages; threads share their parent channel's chain.
|
||||||
|
if message.type == discord.MessageType.default and not isinstance(
|
||||||
|
channel,
|
||||||
|
discord.Thread,
|
||||||
|
):
|
||||||
|
await ingest_message(self.bot.db, message, self.bot.ingest_cache)
|
||||||
|
|
||||||
|
@commands.Cog.listener()
|
||||||
|
async def on_raw_message_delete(
|
||||||
|
self,
|
||||||
|
payload: discord.RawMessageDeleteEvent,
|
||||||
|
) -> None:
|
||||||
|
"""Uningest a recently ingested message when it is deleted.
|
||||||
|
|
||||||
|
:param payload: The raw message delete event.
|
||||||
|
"""
|
||||||
|
await uningest_message(self.bot.db, self.bot.ingest_cache, payload.message_id)
|
||||||
|
|
||||||
|
@commands.Cog.listener()
|
||||||
|
async def on_raw_bulk_message_delete(
|
||||||
|
self,
|
||||||
|
payload: discord.RawBulkMessageDeleteEvent,
|
||||||
|
) -> None:
|
||||||
|
"""Uningest recently ingested messages when they are bulk-deleted.
|
||||||
|
|
||||||
|
:param payload: The raw bulk message delete event.
|
||||||
|
"""
|
||||||
|
for message_id in payload.message_ids:
|
||||||
|
await uningest_message(self.bot.db, self.bot.ingest_cache, message_id)
|
||||||
|
|
||||||
|
|
||||||
|
async def setup(bot: Crabstero) -> None:
|
||||||
|
"""Add the MessageCog to the bot.
|
||||||
|
|
||||||
|
:param bot: The bot instance.
|
||||||
|
"""
|
||||||
|
await bot.add_cog(MessageCog(bot))
|
||||||
@@ -12,11 +12,11 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
"""
|
"""Handles server-level events which trigger channel history ingestion.
|
||||||
Handles server-level events which trigger channel history ingestion.
|
|
||||||
|
|
||||||
All of these events share the same outcome: queuing all textable channels for message history
|
All of these events share the same outcome: queuing all textable channels
|
||||||
ingestion when something changes that may grant new read permissions.
|
for message history ingestion when something changes that may grant new
|
||||||
|
read permissions.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import contextlib
|
import contextlib
|
||||||
@@ -35,13 +35,10 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
|
|
||||||
class ServerEventsCog(commands.Cog):
|
class ServerEventsCog(commands.Cog):
|
||||||
"""
|
"""Cog for handling server-level events that trigger channel history ingestion."""
|
||||||
Cog for handling server-level events that trigger channel history ingestion.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, bot: Crabstero) -> None:
|
def __init__(self, bot: Crabstero) -> None:
|
||||||
"""
|
"""Create a new server events handler cog.
|
||||||
Creates a new server events handler cog.
|
|
||||||
|
|
||||||
:param bot: The bot instance.
|
:param bot: The bot instance.
|
||||||
"""
|
"""
|
||||||
@@ -49,29 +46,30 @@ class ServerEventsCog(commands.Cog):
|
|||||||
|
|
||||||
@commands.Cog.listener()
|
@commands.Cog.listener()
|
||||||
async def on_guild_join(self, guild: discord.Guild) -> None:
|
async def on_guild_join(self, guild: discord.Guild) -> None:
|
||||||
"""
|
"""Queue all text channels for ingestion when joining a new server.
|
||||||
Queues all text channels for message history ingestion when joining a new server.
|
|
||||||
Also logs the join and sends an embed to the bot owner.
|
Also log the join and send an embed to the bot owner.
|
||||||
|
|
||||||
:param guild: The guild that was joined.
|
:param guild: The guild that was joined.
|
||||||
"""
|
"""
|
||||||
logger.info('Joined new server! (Name: "%s" | ID: %s)', guild.name, guild.id)
|
logger.info('Joined new server! (Name: "%s" | ID: %s)', guild.name, guild.id)
|
||||||
|
|
||||||
embed = discord.Embed(
|
embed = discord.Embed(
|
||||||
title="Joined New Server", color=discord.Color.from_rgb(255, 165, 0)
|
title="Joined New Server",
|
||||||
|
color=discord.Color.from_rgb(255, 165, 0),
|
||||||
)
|
)
|
||||||
embed.add_field(
|
embed.add_field(
|
||||||
name=f"{guild.name} ({guild.id})",
|
name=f"{guild.name} ({guild.id})",
|
||||||
value=f"{guild.member_count} members",
|
value=f"{guild.member_count} members",
|
||||||
)
|
)
|
||||||
|
|
||||||
if guild.icon:
|
if guild.icon is not None:
|
||||||
embed.set_image(url=guild.icon.url)
|
embed.set_image(url=guild.icon.url)
|
||||||
|
|
||||||
embed.set_footer(text=f"{len(self.bot.guilds)} total servers")
|
embed.set_footer(text=f"{len(self.bot.guilds)} total servers")
|
||||||
|
|
||||||
app_info = await self.bot.application_info()
|
app_info = await self.bot.application_info()
|
||||||
if app_info.owner:
|
if app_info.owner is not None:
|
||||||
with contextlib.suppress(discord.HTTPException):
|
with contextlib.suppress(discord.HTTPException):
|
||||||
await app_info.owner.send(embed=embed)
|
await app_info.owner.send(embed=embed)
|
||||||
|
|
||||||
@@ -79,8 +77,7 @@ class ServerEventsCog(commands.Cog):
|
|||||||
|
|
||||||
@commands.Cog.listener()
|
@commands.Cog.listener()
|
||||||
async def on_guild_available(self, guild: discord.Guild) -> None:
|
async def on_guild_available(self, guild: discord.Guild) -> None:
|
||||||
"""
|
"""Queue all textable channels for ingestion when a server becomes available.
|
||||||
Queues all textable channels for message history ingestion when a server becomes available.
|
|
||||||
|
|
||||||
:param guild: The guild that became available.
|
:param guild: The guild that became available.
|
||||||
"""
|
"""
|
||||||
@@ -88,11 +85,13 @@ class ServerEventsCog(commands.Cog):
|
|||||||
|
|
||||||
@commands.Cog.listener()
|
@commands.Cog.listener()
|
||||||
async def on_guild_role_update(
|
async def on_guild_role_update(
|
||||||
self, before: discord.Role, after: discord.Role
|
self,
|
||||||
|
before: discord.Role,
|
||||||
|
after: discord.Role,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""Queue all text channels for ingestion when role permissions change.
|
||||||
Queues all text channels for message history ingestion when role permissions change,
|
|
||||||
but only if the bot is a member of the updated role.
|
Only triggers if the bot is a member of the updated role.
|
||||||
|
|
||||||
:param before: The role before the update.
|
:param before: The role before the update.
|
||||||
:param after: The role after the update.
|
:param after: The role after the update.
|
||||||
@@ -105,11 +104,11 @@ class ServerEventsCog(commands.Cog):
|
|||||||
|
|
||||||
@commands.Cog.listener()
|
@commands.Cog.listener()
|
||||||
async def on_guild_channel_update(
|
async def on_guild_channel_update(
|
||||||
self, before: discord.abc.GuildChannel, after: discord.abc.GuildChannel
|
self,
|
||||||
|
before: discord.abc.GuildChannel,
|
||||||
|
after: discord.abc.GuildChannel,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""Queue all text channels for ingestion on permission change.
|
||||||
Queues all text channels for message history ingestion when channel override permissions
|
|
||||||
change.
|
|
||||||
|
|
||||||
:param before: The channel before the update.
|
:param before: The channel before the update.
|
||||||
:param after: The channel after the update.
|
:param after: The channel after the update.
|
||||||
@@ -121,8 +120,7 @@ class ServerEventsCog(commands.Cog):
|
|||||||
|
|
||||||
|
|
||||||
async def setup(bot: Crabstero) -> None:
|
async def setup(bot: Crabstero) -> None:
|
||||||
"""
|
"""Add the ServerEventsCog to the bot.
|
||||||
Adds the ServerEventsCog to the bot.
|
|
||||||
|
|
||||||
:param bot: The bot instance.
|
:param bot: The bot instance.
|
||||||
"""
|
"""
|
||||||
@@ -0,0 +1,215 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Markov chain sentence ingestion and generation backed by SQLite.
|
||||||
|
|
||||||
|
The Markov chain stores word transitions per channel, using occurrence counters
|
||||||
|
to represent frequency weight.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import itertools
|
||||||
|
import re
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from crabstero import metrics
|
||||||
|
from crabstero.database import StartWord, Transition
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from crabstero.database import Database
|
||||||
|
|
||||||
|
# Section sign (§), used internally to mark the end of sentences that lack punctuation,
|
||||||
|
# as is common on Discord.
|
||||||
|
DEFAULT_SENTENCE_END = "\u00a7"
|
||||||
|
|
||||||
|
_TERMINATORS = frozenset({DEFAULT_SENTENCE_END, ".", "!", "?"})
|
||||||
|
_SENTENCE_SPLIT = re.compile(r"(?<=[.!?]) ")
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_whitespace(text: str) -> str:
|
||||||
|
"""Collapse runs of whitespace into single spaces and strip edges.
|
||||||
|
|
||||||
|
:param text: The raw text to normalize.
|
||||||
|
:return: The normalized text.
|
||||||
|
"""
|
||||||
|
return " ".join(text.split())
|
||||||
|
|
||||||
|
|
||||||
|
def is_complete_sentence(sentence: str) -> bool:
|
||||||
|
"""Check whether a sentence ends with a valid terminator.
|
||||||
|
|
||||||
|
Valid terminators are the section sign, period, exclamation mark, or
|
||||||
|
question mark.
|
||||||
|
|
||||||
|
:param sentence: The sentence to test.
|
||||||
|
:return: True if the sentence ends with a valid terminator, False otherwise.
|
||||||
|
"""
|
||||||
|
if not sentence:
|
||||||
|
return False
|
||||||
|
|
||||||
|
return sentence[-1] in _TERMINATORS
|
||||||
|
|
||||||
|
|
||||||
|
def _split_sentences(paragraph: str) -> list[str]:
|
||||||
|
"""Normalize whitespace and split a paragraph into sentences.
|
||||||
|
|
||||||
|
:param paragraph: The raw paragraph text.
|
||||||
|
:return: A list of individual sentences.
|
||||||
|
"""
|
||||||
|
if not is_complete_sentence(paragraph):
|
||||||
|
paragraph += DEFAULT_SENTENCE_END
|
||||||
|
normalized = _normalize_whitespace(paragraph)
|
||||||
|
return _SENTENCE_SPLIT.split(normalized)
|
||||||
|
|
||||||
|
|
||||||
|
_MIN_WORDS_FOR_START = 2
|
||||||
|
|
||||||
|
|
||||||
|
def _tokenize_sentence(
|
||||||
|
sentence: str,
|
||||||
|
) -> tuple[list[str], list[tuple[str, str]]]:
|
||||||
|
"""Tokenize a sentence into start words and transition pairs.
|
||||||
|
|
||||||
|
:param sentence: A single sentence.
|
||||||
|
:return: A tuple of (start_words, transitions) using bare word strings.
|
||||||
|
"""
|
||||||
|
if not is_complete_sentence(sentence):
|
||||||
|
sentence += DEFAULT_SENTENCE_END
|
||||||
|
words = sentence.split()
|
||||||
|
|
||||||
|
start_words: list[str] = [words[0]] if len(words) >= _MIN_WORDS_FOR_START else []
|
||||||
|
transitions: list[tuple[str, str]] = list(itertools.pairwise(words))
|
||||||
|
return start_words, transitions
|
||||||
|
|
||||||
|
|
||||||
|
async def ingest(db: Database, channel_id: int, user_id: int, paragraph: str) -> None:
|
||||||
|
"""Ingest a paragraph into the Markov chain for a given channel.
|
||||||
|
|
||||||
|
The paragraph may contain multiple sentences which are split and ingested
|
||||||
|
individually.
|
||||||
|
|
||||||
|
:param db: The database instance.
|
||||||
|
:param channel_id: The Discord channel ID to associate with this data.
|
||||||
|
:param user_id: The Discord user ID of the contributor.
|
||||||
|
:param paragraph: The paragraph of sentences to ingest.
|
||||||
|
"""
|
||||||
|
for sentence in _split_sentences(paragraph):
|
||||||
|
await _ingest_sentence(db, channel_id, user_id, sentence)
|
||||||
|
|
||||||
|
|
||||||
|
async def _ingest_sentence(
|
||||||
|
db: Database,
|
||||||
|
channel_id: int,
|
||||||
|
user_id: int,
|
||||||
|
sentence: str,
|
||||||
|
) -> None:
|
||||||
|
"""Ingest a single sentence into the Markov chain for a given channel.
|
||||||
|
|
||||||
|
:param db: The database instance.
|
||||||
|
:param channel_id: The Discord channel ID to associate with this data.
|
||||||
|
:param user_id: The Discord user ID of the contributor.
|
||||||
|
:param sentence: The sentence to ingest.
|
||||||
|
"""
|
||||||
|
raw_starts, raw_transitions = _tokenize_sentence(sentence)
|
||||||
|
start_words = [StartWord(channel_id, user_id, w) for w in raw_starts]
|
||||||
|
transitions = [Transition(channel_id, user_id, w, nw) for w, nw in raw_transitions]
|
||||||
|
await db.add_markov_data(start_words, transitions)
|
||||||
|
|
||||||
|
|
||||||
|
async def uningest(db: Database, channel_id: int, user_id: int, paragraph: str) -> None:
|
||||||
|
"""Remove a paragraph's Markov data from the chain for a given channel.
|
||||||
|
|
||||||
|
Mirrors :func:`ingest` but deletes one matching row per entry instead of
|
||||||
|
inserting.
|
||||||
|
|
||||||
|
:param db: The database instance.
|
||||||
|
:param channel_id: The Discord channel ID.
|
||||||
|
:param user_id: The Discord user ID of the contributor.
|
||||||
|
:param paragraph: The paragraph of sentences to uningest.
|
||||||
|
"""
|
||||||
|
for sentence in _split_sentences(paragraph):
|
||||||
|
await _uningest_sentence(db, channel_id, user_id, sentence)
|
||||||
|
|
||||||
|
|
||||||
|
async def _uningest_sentence(
|
||||||
|
db: Database,
|
||||||
|
channel_id: int,
|
||||||
|
user_id: int,
|
||||||
|
sentence: str,
|
||||||
|
) -> None:
|
||||||
|
"""Remove a single sentence's Markov data from the chain.
|
||||||
|
|
||||||
|
:param db: The database instance.
|
||||||
|
:param channel_id: The Discord channel ID.
|
||||||
|
:param user_id: The Discord user ID of the contributor.
|
||||||
|
:param sentence: The sentence to uningest.
|
||||||
|
"""
|
||||||
|
raw_starts, raw_transitions = _tokenize_sentence(sentence)
|
||||||
|
start_words = [StartWord(channel_id, user_id, w) for w in raw_starts]
|
||||||
|
transitions = [Transition(channel_id, user_id, w, nw) for w, nw in raw_transitions]
|
||||||
|
await db.remove_markov_data(start_words, transitions)
|
||||||
|
|
||||||
|
|
||||||
|
async def generate(
|
||||||
|
db: Database,
|
||||||
|
channel_id: int,
|
||||||
|
soft_limit: int = 750,
|
||||||
|
hard_limit: int = 1000,
|
||||||
|
) -> str:
|
||||||
|
"""Generate a new sentence from previously ingested words for a channel.
|
||||||
|
|
||||||
|
:param db: The database instance.
|
||||||
|
:param channel_id: The Discord channel ID to generate from.
|
||||||
|
:param soft_limit: Character count to aim for when wrapping up.
|
||||||
|
:param hard_limit: Character count to hard-cut the sentence at.
|
||||||
|
:return: A new generated sentence.
|
||||||
|
"""
|
||||||
|
word = await db.get_random_start_word(channel_id)
|
||||||
|
|
||||||
|
if word is None:
|
||||||
|
return (
|
||||||
|
"I do not have enough data to generate a message yet."
|
||||||
|
" Chat a bit more so I can learn how this channel talks."
|
||||||
|
)
|
||||||
|
|
||||||
|
with metrics.GENERATION_DURATION.time():
|
||||||
|
parts: list[str] = [word]
|
||||||
|
current_length = len(word)
|
||||||
|
|
||||||
|
# The loop is skipped if the start word already ends a sentence (e.g. "Yes.").
|
||||||
|
while not is_complete_sentence(word):
|
||||||
|
# Past the soft limit, prefer a sentence-ending word to wrap up.
|
||||||
|
if current_length >= soft_limit:
|
||||||
|
next_word = await db.get_random_completing_next_word(channel_id, word)
|
||||||
|
if next_word is None:
|
||||||
|
next_word = await db.get_random_next_word(channel_id, word)
|
||||||
|
else:
|
||||||
|
next_word = await db.get_random_next_word(channel_id, word)
|
||||||
|
|
||||||
|
if next_word is None:
|
||||||
|
break
|
||||||
|
|
||||||
|
word = next_word
|
||||||
|
|
||||||
|
parts.append(word)
|
||||||
|
current_length += 1 + len(word) # +1 for the joining space.
|
||||||
|
|
||||||
|
if current_length >= hard_limit:
|
||||||
|
break
|
||||||
|
|
||||||
|
result = " ".join(parts)
|
||||||
|
if current_length >= hard_limit:
|
||||||
|
result = result[:hard_limit]
|
||||||
|
|
||||||
|
# Strip the internal sentence-end marker so it never appears in output.
|
||||||
|
return result.removesuffix(DEFAULT_SENTENCE_END)
|
||||||
@@ -0,0 +1,233 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Generate Discord reply messages and ingest raw messages.
|
||||||
|
|
||||||
|
Orchestrates Markov chain generation and ingestion in the context of
|
||||||
|
Discord messages, handling reply logic, embed generation, mention
|
||||||
|
filtering, and message ingestion with flag checks.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import re
|
||||||
|
import secrets
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import discord
|
||||||
|
|
||||||
|
from crabstero import flags, markov, metrics
|
||||||
|
from crabstero.cache import CachedMessage, IngestCache
|
||||||
|
from crabstero.database import ChannelImage
|
||||||
|
from crabstero.flags import EntityType, Flag
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from crabstero.database import Database
|
||||||
|
|
||||||
|
# Copied from discordjs/discord-api-types:
|
||||||
|
# https://github.com/discordjs/discord-api-types/blob/662cb0cb0ac9c6f9ad93e180849476714bfceb0c/globals.ts#L39
|
||||||
|
_MENTION_PATTERN = re.compile(r"<@!?(\d{17,20})>")
|
||||||
|
|
||||||
|
_EMBED_CHANCE_THRESHOLD = 95 # Out of 100; sends an embed ~5% of the time.
|
||||||
|
|
||||||
|
|
||||||
|
async def reply_to_message(db: Database, message: discord.Message) -> None:
|
||||||
|
"""Send a new message in Discord in response to a given message.
|
||||||
|
|
||||||
|
:param db: The database instance.
|
||||||
|
:param message: The message prompting the response.
|
||||||
|
"""
|
||||||
|
channel = message.channel
|
||||||
|
guild = message.guild
|
||||||
|
if guild is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
if not channel.permissions_for(guild.me).send_messages:
|
||||||
|
return
|
||||||
|
|
||||||
|
if (
|
||||||
|
await flags.is_flag_set(db, channel, EntityType.CHANNEL, Flag.NO_REPLY)
|
||||||
|
or await flags.is_flag_set(db, guild, EntityType.SERVER, Flag.NO_REPLY)
|
||||||
|
or await flags.is_flag_set(db, message.author, EntityType.USER, Flag.NO_REPLY)
|
||||||
|
):
|
||||||
|
return
|
||||||
|
|
||||||
|
# Threads share their parent channel's Markov chain.
|
||||||
|
if isinstance(channel, discord.Thread):
|
||||||
|
channel_id = channel.parent_id
|
||||||
|
else:
|
||||||
|
channel_id = channel.id
|
||||||
|
|
||||||
|
body = await markov.generate(db, channel_id)
|
||||||
|
|
||||||
|
embed = None
|
||||||
|
|
||||||
|
# 5% chance to include an embed, if the bot has permission.
|
||||||
|
if (
|
||||||
|
secrets.randbelow(100) >= _EMBED_CHANCE_THRESHOLD
|
||||||
|
and channel.permissions_for(guild.me).embed_links
|
||||||
|
):
|
||||||
|
embed = discord.Embed(
|
||||||
|
title=await markov.generate(db, channel_id, soft_limit=200, hard_limit=256),
|
||||||
|
description=await markov.generate(
|
||||||
|
db,
|
||||||
|
channel_id,
|
||||||
|
soft_limit=300,
|
||||||
|
hard_limit=500,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
random_image = await db.get_random_image(channel_id)
|
||||||
|
if random_image is not None:
|
||||||
|
embed.set_image(url=random_image)
|
||||||
|
|
||||||
|
# Suppress all mentions by default; only ping users who opted in.
|
||||||
|
allowed_user_ids: set[int] = set()
|
||||||
|
|
||||||
|
for user_id in {int(uid) for uid in _MENTION_PATTERN.findall(body)}:
|
||||||
|
if await flags.is_flag_set(db, user_id, EntityType.USER, Flag.ALLOW_PINGS):
|
||||||
|
allowed_user_ids.add(user_id)
|
||||||
|
|
||||||
|
allowed_mentions = discord.AllowedMentions(
|
||||||
|
everyone=False,
|
||||||
|
roles=False,
|
||||||
|
users=[discord.Object(id=uid) for uid in allowed_user_ids],
|
||||||
|
)
|
||||||
|
|
||||||
|
if embed is not None:
|
||||||
|
await message.reply(
|
||||||
|
content=body,
|
||||||
|
embed=embed,
|
||||||
|
allowed_mentions=allowed_mentions,
|
||||||
|
mention_author=False,
|
||||||
|
)
|
||||||
|
metrics.EMBEDS_GENERATED.inc()
|
||||||
|
else:
|
||||||
|
await message.reply(
|
||||||
|
content=body,
|
||||||
|
allowed_mentions=allowed_mentions,
|
||||||
|
mention_author=False,
|
||||||
|
)
|
||||||
|
metrics.REPLIES_SENT.inc()
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_embed_data(
|
||||||
|
embeds: list[discord.Embed],
|
||||||
|
) -> tuple[list[str], list[str]]:
|
||||||
|
"""Extract text and image URLs from message embeds.
|
||||||
|
|
||||||
|
:param embeds: The embeds to process.
|
||||||
|
:return: A tuple of (texts, image_urls).
|
||||||
|
"""
|
||||||
|
texts: list[str] = []
|
||||||
|
image_urls: list[str] = []
|
||||||
|
for embed in embeds:
|
||||||
|
if embed.title:
|
||||||
|
texts.append(embed.title)
|
||||||
|
if embed.description:
|
||||||
|
texts.append(embed.description)
|
||||||
|
if embed.image and embed.image.url:
|
||||||
|
image_urls.append(embed.image.url)
|
||||||
|
return texts, image_urls
|
||||||
|
|
||||||
|
|
||||||
|
async def ingest_message(
|
||||||
|
db: Database,
|
||||||
|
message: discord.Message,
|
||||||
|
cache: IngestCache | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Ingest a given message into its channel's Markov chain.
|
||||||
|
|
||||||
|
If a cache is provided, the ingested data is recorded so it can be
|
||||||
|
reversed by :func:`uningest_message` if the message is deleted shortly
|
||||||
|
after.
|
||||||
|
|
||||||
|
:param db: The database instance.
|
||||||
|
:param message: The message to ingest.
|
||||||
|
:param cache: Optional ingest cache for tracking recently ingested data.
|
||||||
|
"""
|
||||||
|
guild = message.guild
|
||||||
|
if guild is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
if message.author.bot or guild.me in message.mentions:
|
||||||
|
return
|
||||||
|
|
||||||
|
if (
|
||||||
|
await flags.is_flag_set(db, message.channel, EntityType.CHANNEL, Flag.NO_INGEST)
|
||||||
|
or await flags.is_flag_set(db, guild, EntityType.SERVER, Flag.NO_INGEST)
|
||||||
|
or await flags.is_flag_set(db, message.author, EntityType.USER, Flag.NO_INGEST)
|
||||||
|
):
|
||||||
|
return
|
||||||
|
|
||||||
|
channel_id = message.channel.id
|
||||||
|
user_id = message.author.id
|
||||||
|
|
||||||
|
if not message.content and not message.embeds:
|
||||||
|
return
|
||||||
|
|
||||||
|
embed_texts, image_urls = _extract_embed_data(message.embeds)
|
||||||
|
|
||||||
|
with metrics.MESSAGE_INGESTION_DURATION.time():
|
||||||
|
if message.content:
|
||||||
|
await markov.ingest(db, channel_id, user_id, message.content)
|
||||||
|
|
||||||
|
for text in embed_texts:
|
||||||
|
await markov.ingest(db, channel_id, user_id, text)
|
||||||
|
|
||||||
|
if image_urls:
|
||||||
|
await db.add_images(
|
||||||
|
[ChannelImage(channel_id, user_id, url) for url in image_urls],
|
||||||
|
)
|
||||||
|
|
||||||
|
if cache is not None:
|
||||||
|
cache.put(
|
||||||
|
message.id,
|
||||||
|
CachedMessage(
|
||||||
|
channel_id=channel_id,
|
||||||
|
user_id=user_id,
|
||||||
|
content=message.content or None,
|
||||||
|
embed_texts=embed_texts,
|
||||||
|
image_urls=image_urls,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def uningest_message(db: Database, cache: IngestCache, message_id: int) -> None:
|
||||||
|
"""Reverse ingestion for a recently deleted message.
|
||||||
|
|
||||||
|
Looks up the message in the cache. If found, removes all Markov data
|
||||||
|
and images that were added during ingestion.
|
||||||
|
|
||||||
|
:param db: The database instance.
|
||||||
|
:param cache: The ingest cache.
|
||||||
|
:param message_id: The Discord message ID that was deleted.
|
||||||
|
"""
|
||||||
|
entry = cache.pop(message_id)
|
||||||
|
if entry is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
if entry.content:
|
||||||
|
await markov.uningest(db, entry.channel_id, entry.user_id, entry.content)
|
||||||
|
|
||||||
|
for text in entry.embed_texts:
|
||||||
|
await markov.uningest(db, entry.channel_id, entry.user_id, text)
|
||||||
|
|
||||||
|
if entry.image_urls:
|
||||||
|
await db.remove_images(
|
||||||
|
[
|
||||||
|
ChannelImage(entry.channel_id, entry.user_id, url)
|
||||||
|
for url in entry.image_urls
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
metrics.MESSAGES_UNINGESTED.inc()
|
||||||
@@ -0,0 +1,257 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Prometheus metrics definitions and HTTP server for Crabstero.
|
||||||
|
|
||||||
|
All metric objects are defined at module level using the default global
|
||||||
|
registry. The MetricsServer class wraps an aiohttp application that serves
|
||||||
|
the ``/metrics`` scrape endpoint.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
import errno
|
||||||
|
import logging
|
||||||
|
import socket
|
||||||
|
import stat
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from aiohttp import web
|
||||||
|
from prometheus_client import Counter, Gauge, Histogram, Info
|
||||||
|
from prometheus_client.aiohttp import make_aiohttp_handler
|
||||||
|
|
||||||
|
from crabstero import __version__
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_FAST_BUCKETS = (0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0)
|
||||||
|
_SLOW_BUCKETS = (0.1, 0.5, 1.0, 5.0, 10.0, 30.0, 60.0, 120.0, 300.0, 600.0)
|
||||||
|
_MESSAGES_BUCKETS = (100, 500, 1000, 5000, 10000, 25000, 50000)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class TcpMetricsAddress:
|
||||||
|
"""TCP listen address for the metrics server."""
|
||||||
|
|
||||||
|
host: str
|
||||||
|
port: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class UnixMetricsAddress:
|
||||||
|
"""Unix socket listen address for the metrics server."""
|
||||||
|
|
||||||
|
path: str
|
||||||
|
mode: int | None = None
|
||||||
|
|
||||||
|
|
||||||
|
type MetricsAddress = TcpMetricsAddress | UnixMetricsAddress
|
||||||
|
|
||||||
|
|
||||||
|
BUILD_INFO = Info("crabstero_build", "Build information")
|
||||||
|
BUILD_INFO.info({"version": __version__})
|
||||||
|
|
||||||
|
MESSAGES_UNINGESTED = Counter(
|
||||||
|
"crabstero_messages_uningested_total",
|
||||||
|
"Messages uningested from a Markov chain after deletion",
|
||||||
|
)
|
||||||
|
REPLIES_SENT = Counter(
|
||||||
|
"crabstero_replies_sent_total",
|
||||||
|
"Markov replies sent",
|
||||||
|
)
|
||||||
|
EMBEDS_GENERATED = Counter(
|
||||||
|
"crabstero_embeds_generated_total",
|
||||||
|
"Embeds included in replies",
|
||||||
|
)
|
||||||
|
MESSAGES_PROCESSED = Counter(
|
||||||
|
"crabstero_messages_processed_total",
|
||||||
|
"Messages seen by the listener",
|
||||||
|
)
|
||||||
|
PINGME = Counter(
|
||||||
|
"crabstero_pingme_total",
|
||||||
|
"/pingme command outcomes",
|
||||||
|
["outcome"],
|
||||||
|
)
|
||||||
|
PINGME.labels(outcome="opted_in")
|
||||||
|
PINGME.labels(outcome="opted_out")
|
||||||
|
|
||||||
|
FORGETME = Counter(
|
||||||
|
"crabstero_forgetme_total",
|
||||||
|
"/forgetme command outcomes",
|
||||||
|
["outcome"],
|
||||||
|
)
|
||||||
|
FORGETME.labels(outcome="completed")
|
||||||
|
FORGETME.labels(outcome="cancelled")
|
||||||
|
FORGETME.labels(outcome="timeout")
|
||||||
|
FORGETME.labels(outcome="already_forgotten")
|
||||||
|
|
||||||
|
ERRORS = Counter(
|
||||||
|
"crabstero_errors_total",
|
||||||
|
"Unhandled exceptions by source",
|
||||||
|
["source"],
|
||||||
|
)
|
||||||
|
DISCORD_EVENTS = Counter(
|
||||||
|
"crabstero_discord_events_total",
|
||||||
|
"Discord gateway events received",
|
||||||
|
["event"],
|
||||||
|
)
|
||||||
|
|
||||||
|
DISCORD_LATENCY = Gauge(
|
||||||
|
"crabstero_discord_latency_seconds",
|
||||||
|
"Discord WebSocket heartbeat latency",
|
||||||
|
)
|
||||||
|
GUILD_COUNT = Gauge(
|
||||||
|
"crabstero_guild_count",
|
||||||
|
"Guilds the bot is currently in",
|
||||||
|
)
|
||||||
|
INGESTION_ACTIVE = Gauge(
|
||||||
|
"crabstero_ingestion_active_channels",
|
||||||
|
"Channels pending or in-progress for ingestion",
|
||||||
|
)
|
||||||
|
|
||||||
|
MESSAGE_INGESTION_DURATION = Histogram(
|
||||||
|
"crabstero_message_ingestion_duration_seconds",
|
||||||
|
"Time to ingest one message",
|
||||||
|
buckets=_FAST_BUCKETS,
|
||||||
|
)
|
||||||
|
CHANNEL_INGESTION_DURATION = Histogram(
|
||||||
|
"crabstero_channel_ingestion_duration_seconds",
|
||||||
|
"Time to fully ingest a channel",
|
||||||
|
buckets=_SLOW_BUCKETS,
|
||||||
|
)
|
||||||
|
GENERATION_DURATION = Histogram(
|
||||||
|
"crabstero_generation_duration_seconds",
|
||||||
|
"Time to generate a Markov response",
|
||||||
|
buckets=_FAST_BUCKETS,
|
||||||
|
)
|
||||||
|
CHANNEL_INGESTION_MESSAGES = Histogram(
|
||||||
|
"crabstero_channel_ingestion_messages",
|
||||||
|
"Messages iterated per channel ingestion run",
|
||||||
|
buckets=_MESSAGES_BUCKETS,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class MetricsServer:
|
||||||
|
"""HTTP server that exposes a Prometheus ``/metrics`` scrape endpoint.
|
||||||
|
|
||||||
|
Uses ``aiohttp.web.AppRunner`` with TCP or Unix socket sites for
|
||||||
|
async-native serving.
|
||||||
|
The handler is provided by ``prometheus_client.aiohttp.make_aiohttp_handler``,
|
||||||
|
which handles compression and content negotiation automatically.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("_address", "_port", "_runner")
|
||||||
|
|
||||||
|
def __init__(self, address: MetricsAddress) -> None:
|
||||||
|
"""Store the listen address.
|
||||||
|
|
||||||
|
:param address: TCP or Unix socket address to bind to.
|
||||||
|
"""
|
||||||
|
self._address = address
|
||||||
|
self._port = address.port if isinstance(address, TcpMetricsAddress) else None
|
||||||
|
self._runner: web.AppRunner | None = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def port(self) -> int:
|
||||||
|
"""The TCP port the server is bound to.
|
||||||
|
|
||||||
|
After ``start()``, this reflects the actual port (useful when
|
||||||
|
the constructor received port 0 for OS assignment).
|
||||||
|
"""
|
||||||
|
if self._port is None:
|
||||||
|
msg = "Unix socket metrics server does not have a TCP port"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
return self._port
|
||||||
|
|
||||||
|
async def start(self) -> None:
|
||||||
|
"""Create the aiohttp application and start listening."""
|
||||||
|
if self._runner is not None:
|
||||||
|
msg = "Metrics server is already running"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
app = web.Application()
|
||||||
|
app.router.add_get("/metrics", make_aiohttp_handler())
|
||||||
|
self._runner = web.AppRunner(app)
|
||||||
|
await self._runner.setup()
|
||||||
|
try:
|
||||||
|
if isinstance(self._address, TcpMetricsAddress):
|
||||||
|
site: web.BaseSite = web.TCPSite(
|
||||||
|
self._runner,
|
||||||
|
self._address.host,
|
||||||
|
self._address.port,
|
||||||
|
)
|
||||||
|
await site.start()
|
||||||
|
|
||||||
|
# Resolve the actual bound port when the OS assigned one.
|
||||||
|
if self._address.port == 0:
|
||||||
|
self._port = self._runner.addresses[0][1]
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Metrics server listening on %s:%d.",
|
||||||
|
self._address.host,
|
||||||
|
self.port,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
socket_path = Path(self._address.path)
|
||||||
|
# This runs once before binding the Unix site, so synchronous
|
||||||
|
# path checks are simpler than offloading startup setup.
|
||||||
|
try:
|
||||||
|
mode = socket_path.stat( # noqa: ASYNC240
|
||||||
|
follow_symlinks=False,
|
||||||
|
).st_mode
|
||||||
|
except FileNotFoundError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
if not stat.S_ISSOCK(mode):
|
||||||
|
msg = (
|
||||||
|
"Metrics Unix socket path already exists and is not a "
|
||||||
|
f"socket: {self._address.path}"
|
||||||
|
)
|
||||||
|
raise FileExistsError(msg)
|
||||||
|
|
||||||
|
with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as probe:
|
||||||
|
probe.settimeout(0.1)
|
||||||
|
try:
|
||||||
|
probe.connect(self._address.path)
|
||||||
|
except OSError as exc:
|
||||||
|
if exc.errno not in {errno.ECONNREFUSED, errno.ENOENT}:
|
||||||
|
raise
|
||||||
|
else:
|
||||||
|
msg = (
|
||||||
|
"Metrics Unix socket path is already in use: "
|
||||||
|
f"{self._address.path}"
|
||||||
|
)
|
||||||
|
raise OSError(errno.EADDRINUSE, msg)
|
||||||
|
|
||||||
|
with contextlib.suppress(FileNotFoundError):
|
||||||
|
socket_path.unlink() # noqa: ASYNC240
|
||||||
|
|
||||||
|
site = web.UnixSite(self._runner, self._address.path)
|
||||||
|
await site.start()
|
||||||
|
if self._address.mode is not None:
|
||||||
|
socket_path.chmod(self._address.mode) # noqa: ASYNC240
|
||||||
|
logger.info(
|
||||||
|
"Metrics server listening on Unix socket %s.",
|
||||||
|
self._address.path,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
await self._runner.cleanup()
|
||||||
|
self._runner = None
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
"""Shut down the HTTP server and release resources."""
|
||||||
|
if self._runner is not None:
|
||||||
|
await self._runner.cleanup()
|
||||||
|
self._runner = None
|
||||||
|
logger.info("Metrics server stopped.")
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Background tasks for Crabstero."""
|
||||||
@@ -12,10 +12,9 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
"""
|
"""Bulk-ingests the message history of channels.
|
||||||
Bulk-ingests the message history of channels.
|
|
||||||
|
|
||||||
Channels are enqueued for processing by the bot's fixed ingestion worker pool.
|
Channels are processed by the bot's semaphore-bounded dynamic tasks.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
@@ -23,10 +22,11 @@ from typing import TYPE_CHECKING
|
|||||||
|
|
||||||
import discord
|
import discord
|
||||||
|
|
||||||
|
from crabstero import metrics
|
||||||
from crabstero.messages import ingest_message
|
from crabstero.messages import ingest_message
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from crabstero.bot import Crabstero
|
from crabstero.bot import Crabstero, IngestableChannel
|
||||||
from crabstero.database import Database
|
from crabstero.database import Database
|
||||||
|
|
||||||
MAXIMUM_MESSAGES_PER_CHANNEL = (
|
MAXIMUM_MESSAGES_PER_CHANNEL = (
|
||||||
@@ -37,8 +37,7 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
|
|
||||||
def queue_channels_for_ingestion(guild: discord.Guild, bot: Crabstero) -> None:
|
def queue_channels_for_ingestion(guild: discord.Guild, bot: Crabstero) -> None:
|
||||||
"""
|
"""Enqueue every textable channel in a guild for message history ingestion.
|
||||||
Enqueues every textable channel in a guild for message history ingestion.
|
|
||||||
|
|
||||||
:param guild: The Discord guild whose channels should be ingested.
|
:param guild: The Discord guild whose channels should be ingested.
|
||||||
:param bot: The bot instance.
|
:param bot: The bot instance.
|
||||||
@@ -49,20 +48,21 @@ def queue_channels_for_ingestion(guild: discord.Guild, bot: Crabstero) -> None:
|
|||||||
|
|
||||||
|
|
||||||
async def ingest_channel(
|
async def ingest_channel(
|
||||||
channel: discord.TextChannel | discord.VoiceChannel,
|
channel: IngestableChannel,
|
||||||
db: Database,
|
db: Database,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""Bulk-ingest the message history of a given channel.
|
||||||
Bulk-ingests the message history of a given channel. The task will be ended early if
|
|
||||||
permissions do not allow ingesting this channel or if it has already been ingested in the past.
|
End early if permissions do not allow ingesting this channel or if it has
|
||||||
|
already been ingested.
|
||||||
|
|
||||||
:param channel: The channel to ingest.
|
:param channel: The channel to ingest.
|
||||||
:param db: The database instance.
|
:param db: The database instance.
|
||||||
"""
|
"""
|
||||||
try:
|
|
||||||
if not channel.permissions_for(channel.guild.me).read_message_history:
|
if not channel.permissions_for(channel.guild.me).read_message_history:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"[%s] Unable to ingest textable channel history due to lacking permissions. Ignoring.",
|
"[%s] Unable to ingest channel history"
|
||||||
|
" due to lacking permissions. Ignoring.",
|
||||||
channel.id,
|
channel.id,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
@@ -74,19 +74,16 @@ async def ingest_channel(
|
|||||||
|
|
||||||
logger.info("[%s] Starting ingestion of textable channel history.", channel.id)
|
logger.info("[%s] Starting ingestion of textable channel history.", channel.id)
|
||||||
|
|
||||||
|
with metrics.CHANNEL_INGESTION_DURATION.time():
|
||||||
count = 0
|
count = 0
|
||||||
async for message in channel.history(limit=MAXIMUM_MESSAGES_PER_CHANNEL):
|
async for message in channel.history(limit=MAXIMUM_MESSAGES_PER_CHANNEL):
|
||||||
count += 1
|
count += 1
|
||||||
await ingest_message(db, message)
|
await ingest_message(db, message)
|
||||||
|
|
||||||
|
metrics.CHANNEL_INGESTION_MESSAGES.observe(count)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"[%s] Ingestion of textable channel history complete. %d messages ingested.",
|
"[%s] Ingestion of channel history complete. %d messages iterated.",
|
||||||
channel.id,
|
channel.id,
|
||||||
count,
|
count,
|
||||||
)
|
)
|
||||||
|
|
||||||
except Exception:
|
|
||||||
logger.exception(
|
|
||||||
"[%s] An error occurred while ingesting textable channel history!",
|
|
||||||
channel.id,
|
|
||||||
)
|
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Test suite for Crabstero."""
|
||||||
|
|||||||
@@ -11,3 +11,37 @@
|
|||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Shared pytest hooks for the test suite."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.hookimpl(tryfirst=True)
|
||||||
|
def pytest_collection_modifyitems(
|
||||||
|
config: pytest.Config,
|
||||||
|
items: list[pytest.Item],
|
||||||
|
) -> None:
|
||||||
|
"""Apply and enforce primary test category markers by directory."""
|
||||||
|
tests_root = config.rootpath / "tests"
|
||||||
|
category_roots = (
|
||||||
|
(tests_root / "unit", pytest.mark.unit),
|
||||||
|
(tests_root / "integration", pytest.mark.integration),
|
||||||
|
)
|
||||||
|
uncategorized: list[str] = []
|
||||||
|
|
||||||
|
for item in items:
|
||||||
|
for category_root, marker in category_roots:
|
||||||
|
if item.path.is_relative_to(category_root):
|
||||||
|
item.add_marker(marker)
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
uncategorized.append(str(item.path.relative_to(config.rootpath)))
|
||||||
|
|
||||||
|
if uncategorized:
|
||||||
|
formatted_paths = "\n".join(f" - {path}" for path in uncategorized)
|
||||||
|
msg = (
|
||||||
|
"Tests must live under tests/unit or tests/integration.\n"
|
||||||
|
f"Uncategorized tests:\n{formatted_paths}"
|
||||||
|
)
|
||||||
|
raise pytest.UsageError(msg)
|
||||||
|
|||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Integration tests for Crabstero."""
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""CLI integration tests for Crabstero."""
|
||||||
@@ -0,0 +1,89 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Integration tests for systemd notification socket behavior."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import socket
|
||||||
|
import sys
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero.cli import _sd_notify
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.skipif(
|
||||||
|
os.name != "posix" or not hasattr(socket, "AF_UNIX"),
|
||||||
|
reason="systemd notification sockets require Unix-domain socket support",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestSystemdNotifySocket:
|
||||||
|
"""Systemd notification helper behavior against real Unix-domain sockets."""
|
||||||
|
|
||||||
|
def test_sd_notify_sends_datagram(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""_sd_notify sends the payload to NOTIFY_SOCKET."""
|
||||||
|
socket_path = tmp_path / "notify.sock"
|
||||||
|
|
||||||
|
with socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM) as server:
|
||||||
|
server.bind(str(socket_path))
|
||||||
|
server.settimeout(1)
|
||||||
|
monkeypatch.setenv("NOTIFY_SOCKET", str(socket_path))
|
||||||
|
|
||||||
|
_sd_notify("READY=1")
|
||||||
|
|
||||||
|
assert server.recv(1024) == b"READY=1"
|
||||||
|
|
||||||
|
def test_sd_notify_ignores_socket_errors(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
caplog: pytest.LogCaptureFixture,
|
||||||
|
) -> None:
|
||||||
|
"""_sd_notify logs and suppresses notification socket send failures."""
|
||||||
|
monkeypatch.setenv("NOTIFY_SOCKET", str(tmp_path / "missing.sock"))
|
||||||
|
|
||||||
|
with caplog.at_level(logging.DEBUG, logger="crabstero"):
|
||||||
|
_sd_notify("READY=1")
|
||||||
|
|
||||||
|
assert "Could not send systemd notification" in caplog.text
|
||||||
|
|
||||||
|
@pytest.mark.skipif(
|
||||||
|
sys.platform != "linux",
|
||||||
|
reason="abstract Unix sockets are Linux-specific",
|
||||||
|
)
|
||||||
|
def test_sd_notify_sends_to_abstract_socket(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""@-prefixed NOTIFY_SOCKET values target Linux abstract sockets."""
|
||||||
|
socket_name = f"crabstero-notify-{os.getpid()}"
|
||||||
|
|
||||||
|
with socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM) as server:
|
||||||
|
server.bind(f"\0{socket_name}")
|
||||||
|
server.settimeout(1)
|
||||||
|
monkeypatch.setenv("NOTIFY_SOCKET", f"@{socket_name}")
|
||||||
|
|
||||||
|
_sd_notify("READY=1")
|
||||||
|
|
||||||
|
assert server.recv(1024) == b"READY=1"
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Database integration tests for Crabstero."""
|
||||||
@@ -0,0 +1,35 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Fixtures used only by database-bound integration tests."""
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero.database import Database
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import AsyncGenerator
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
async def db(tmp_path: Path) -> AsyncGenerator[Database]:
|
||||||
|
"""Yield an isolated file-backed database for SQLite integration tests."""
|
||||||
|
database = await Database.connect(str(tmp_path / "crabstero.db"))
|
||||||
|
try:
|
||||||
|
yield database
|
||||||
|
finally:
|
||||||
|
await database.close()
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,297 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Integration tests for the full ingest → uningest cycle."""
|
||||||
|
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import aiosqlite
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero.cache import CachedMessage, IngestCache
|
||||||
|
from crabstero.database import ChannelImage
|
||||||
|
from crabstero.markov import ingest, uningest
|
||||||
|
from crabstero.messages import uningest_message
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from crabstero.database import Database
|
||||||
|
|
||||||
|
type SnapshotRows = Callable[["Database"], Awaitable[dict[str, list[aiosqlite.Row]]]]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def snapshot_rows() -> SnapshotRows:
|
||||||
|
"""Return a snapshot reader for Markov and image tables."""
|
||||||
|
|
||||||
|
async def read(db: Database) -> dict[str, list[aiosqlite.Row]]:
|
||||||
|
tables: dict[str, list[aiosqlite.Row]] = {}
|
||||||
|
for table in ("markov_start_words", "markov_transitions", "channel_images"):
|
||||||
|
async with db._connection.execute(
|
||||||
|
f"SELECT * FROM {table}", # noqa: S608
|
||||||
|
) as cursor:
|
||||||
|
tables[table] = sorted(await cursor.fetchall())
|
||||||
|
return tables
|
||||||
|
|
||||||
|
return read
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def ingest_cache() -> IngestCache:
|
||||||
|
"""Return an empty ingest cache for uningest-message tests."""
|
||||||
|
return IngestCache()
|
||||||
|
|
||||||
|
|
||||||
|
class TestIngestUningestCycle:
|
||||||
|
"""Full round-trip: ingest data, then uningest it completely."""
|
||||||
|
|
||||||
|
async def test_content_round_trip(self, db: Database) -> None:
|
||||||
|
"""Ingest and uningest content leaves the database clean."""
|
||||||
|
await ingest(db, 1, 100, "Hello beautiful world.")
|
||||||
|
await uningest(db, 1, 100, "Hello beautiful world.")
|
||||||
|
|
||||||
|
assert await db.get_random_start_word(1) is None
|
||||||
|
assert await db.get_random_next_word(1, "Hello") is None
|
||||||
|
assert await db.get_random_next_word(1, "beautiful") is None
|
||||||
|
|
||||||
|
async def test_image_round_trip(self, db: Database) -> None:
|
||||||
|
"""Ingest and uningest an image leaves the database clean."""
|
||||||
|
await db.add_images([ChannelImage(1, 100, "https://example.com/cat.png")])
|
||||||
|
await db.remove_images([ChannelImage(1, 100, "https://example.com/cat.png")])
|
||||||
|
|
||||||
|
assert await db.get_random_image(1) is None
|
||||||
|
|
||||||
|
async def test_multi_sentence_round_trip(self, db: Database) -> None:
|
||||||
|
"""Multi-sentence ingest and uningest leaves the database clean."""
|
||||||
|
text = "Hello world. Goodbye world! How are you?"
|
||||||
|
await ingest(db, 1, 100, text)
|
||||||
|
await uningest(db, 1, 100, text)
|
||||||
|
|
||||||
|
assert await db.get_random_start_word(1) is None
|
||||||
|
|
||||||
|
async def test_uningest_only_removes_one_copy(self, db: Database) -> None:
|
||||||
|
"""Uningesting once preserves data from a second identical ingest."""
|
||||||
|
await ingest(db, 1, 100, "Hello world.")
|
||||||
|
await ingest(db, 1, 100, "Hello world.")
|
||||||
|
await uningest(db, 1, 100, "Hello world.")
|
||||||
|
|
||||||
|
# One copy of each row should remain.
|
||||||
|
assert await db.get_random_start_word(1) == "Hello"
|
||||||
|
assert await db.get_random_next_word(1, "Hello") == "world."
|
||||||
|
|
||||||
|
async def test_uningest_does_not_affect_other_channels(self, db: Database) -> None:
|
||||||
|
"""Uningesting from one channel leaves another channel's data intact."""
|
||||||
|
await ingest(db, 1, 100, "Hello world.")
|
||||||
|
await ingest(db, 2, 100, "Hello world.")
|
||||||
|
await uningest(db, 1, 100, "Hello world.")
|
||||||
|
|
||||||
|
assert await db.get_random_start_word(1) is None
|
||||||
|
assert await db.get_random_start_word(2) == "Hello"
|
||||||
|
assert await db.get_random_next_word(2, "Hello") == "world."
|
||||||
|
|
||||||
|
async def test_uningest_does_not_affect_other_users(self, db: Database) -> None:
|
||||||
|
"""Uningesting for one user leaves another user's data."""
|
||||||
|
await ingest(db, 1, 100, "Hello world.")
|
||||||
|
await ingest(db, 1, 200, "Hello world.")
|
||||||
|
await uningest(db, 1, 100, "Hello world.")
|
||||||
|
|
||||||
|
assert await db.get_random_start_word(1) == "Hello"
|
||||||
|
assert await db.get_random_next_word(1, "Hello") == "world."
|
||||||
|
|
||||||
|
|
||||||
|
class TestUningestRestoresState:
|
||||||
|
"""Uningest restores the database to its prior state."""
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"text",
|
||||||
|
[
|
||||||
|
pytest.param("Hello world.", id="simple-sentence"),
|
||||||
|
pytest.param("Hello world", id="missing-punctuation"),
|
||||||
|
pytest.param(
|
||||||
|
"Hello world. Goodbye world! How are you?",
|
||||||
|
id="multi-sentence",
|
||||||
|
),
|
||||||
|
pytest.param("One.", id="single-word"),
|
||||||
|
pytest.param(
|
||||||
|
"Lots of extra spaces\nand\nnewlines here.",
|
||||||
|
id="whitespace-normalization",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_uningest_restores_empty_db(
|
||||||
|
self,
|
||||||
|
db: Database,
|
||||||
|
text: str,
|
||||||
|
snapshot_rows: SnapshotRows,
|
||||||
|
) -> None:
|
||||||
|
"""Ingest then uningest on an empty database leaves all tables empty."""
|
||||||
|
before = await snapshot_rows(db)
|
||||||
|
await ingest(db, 1, 100, text)
|
||||||
|
await uningest(db, 1, 100, text)
|
||||||
|
after = await snapshot_rows(db)
|
||||||
|
assert after == before
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"text",
|
||||||
|
[
|
||||||
|
pytest.param("Hello world.", id="simple-sentence"),
|
||||||
|
pytest.param(
|
||||||
|
"Hello world. Goodbye world! How are you?",
|
||||||
|
id="multi-sentence",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_uningest_restores_preexisting_data(
|
||||||
|
self,
|
||||||
|
db: Database,
|
||||||
|
text: str,
|
||||||
|
snapshot_rows: SnapshotRows,
|
||||||
|
) -> None:
|
||||||
|
"""Ingest then uningest preserves unrelated pre-existing data exactly."""
|
||||||
|
await ingest(db, 99, 200, "Pre-existing data stays safe.")
|
||||||
|
await db.add_images([ChannelImage(99, 200, "https://example.com/existing.png")])
|
||||||
|
|
||||||
|
before = await snapshot_rows(db)
|
||||||
|
await ingest(db, 1, 100, text)
|
||||||
|
await uningest(db, 1, 100, text)
|
||||||
|
after = await snapshot_rows(db)
|
||||||
|
assert after == before
|
||||||
|
|
||||||
|
|
||||||
|
class TestUningestMessage:
|
||||||
|
"""Orchestrated uningest via cache lookup and database reversal."""
|
||||||
|
|
||||||
|
async def test_content_only(self, db: Database, ingest_cache: IngestCache) -> None:
|
||||||
|
"""Uningest reverses a content-only message via the cache."""
|
||||||
|
await ingest(db, 1, 100, "Hello beautiful world.")
|
||||||
|
ingest_cache.put(
|
||||||
|
555,
|
||||||
|
CachedMessage(
|
||||||
|
channel_id=1,
|
||||||
|
user_id=100,
|
||||||
|
content="Hello beautiful world.",
|
||||||
|
embed_texts=[],
|
||||||
|
image_urls=[],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
await uningest_message(db, ingest_cache, 555)
|
||||||
|
|
||||||
|
assert await db.get_random_start_word(1) is None
|
||||||
|
assert await db.get_random_next_word(1, "Hello") is None
|
||||||
|
|
||||||
|
async def test_embeds_only(self, db: Database, ingest_cache: IngestCache) -> None:
|
||||||
|
"""Uningest reverses embed text ingestion."""
|
||||||
|
await ingest(db, 1, 100, "Embed title here.")
|
||||||
|
await ingest(db, 1, 100, "Embed description here.")
|
||||||
|
ingest_cache.put(
|
||||||
|
556,
|
||||||
|
CachedMessage(
|
||||||
|
channel_id=1,
|
||||||
|
user_id=100,
|
||||||
|
content=None,
|
||||||
|
embed_texts=["Embed title here.", "Embed description here."],
|
||||||
|
image_urls=[],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
await uningest_message(db, ingest_cache, 556)
|
||||||
|
|
||||||
|
assert await db.get_random_start_word(1) is None
|
||||||
|
|
||||||
|
async def test_content_with_embeds_and_images(
|
||||||
|
self,
|
||||||
|
db: Database,
|
||||||
|
ingest_cache: IngestCache,
|
||||||
|
) -> None:
|
||||||
|
"""Uningest reverses content, embed text, and image data together."""
|
||||||
|
await ingest(db, 1, 100, "Body text here.")
|
||||||
|
await ingest(db, 1, 100, "Embed title.")
|
||||||
|
await db.add_images([ChannelImage(1, 100, "https://example.com/img.png")])
|
||||||
|
ingest_cache.put(
|
||||||
|
557,
|
||||||
|
CachedMessage(
|
||||||
|
channel_id=1,
|
||||||
|
user_id=100,
|
||||||
|
content="Body text here.",
|
||||||
|
embed_texts=["Embed title."],
|
||||||
|
image_urls=["https://example.com/img.png"],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
await uningest_message(db, ingest_cache, 557)
|
||||||
|
|
||||||
|
assert await db.get_random_start_word(1) is None
|
||||||
|
assert await db.get_random_image(1) is None
|
||||||
|
|
||||||
|
async def test_cache_miss_is_noop(
|
||||||
|
self,
|
||||||
|
db: Database,
|
||||||
|
ingest_cache: IngestCache,
|
||||||
|
snapshot_rows: SnapshotRows,
|
||||||
|
) -> None:
|
||||||
|
"""A message not in the cache leaves the database unchanged."""
|
||||||
|
await ingest(db, 1, 100, "Keep this data.")
|
||||||
|
before = await snapshot_rows(db)
|
||||||
|
|
||||||
|
await uningest_message(db, ingest_cache, 999)
|
||||||
|
|
||||||
|
after = await snapshot_rows(db)
|
||||||
|
assert after == before
|
||||||
|
|
||||||
|
async def test_preserves_other_messages(
|
||||||
|
self,
|
||||||
|
db: Database,
|
||||||
|
ingest_cache: IngestCache,
|
||||||
|
) -> None:
|
||||||
|
"""Uningesting one message leaves another message's data intact."""
|
||||||
|
await ingest(db, 1, 100, "First message.")
|
||||||
|
await ingest(db, 1, 100, "Second message.")
|
||||||
|
ingest_cache.put(
|
||||||
|
601,
|
||||||
|
CachedMessage(
|
||||||
|
channel_id=1,
|
||||||
|
user_id=100,
|
||||||
|
content="First message.",
|
||||||
|
embed_texts=[],
|
||||||
|
image_urls=[],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
await uningest_message(db, ingest_cache, 601)
|
||||||
|
|
||||||
|
assert await db.get_random_start_word(1) == "Second"
|
||||||
|
assert await db.get_random_next_word(1, "Second") == "message."
|
||||||
|
|
||||||
|
async def test_pops_entry_from_cache(
|
||||||
|
self,
|
||||||
|
db: Database,
|
||||||
|
ingest_cache: IngestCache,
|
||||||
|
) -> None:
|
||||||
|
"""The cache entry is consumed after uningest."""
|
||||||
|
await ingest(db, 1, 100, "Hello world.")
|
||||||
|
ingest_cache.put(
|
||||||
|
602,
|
||||||
|
CachedMessage(
|
||||||
|
channel_id=1,
|
||||||
|
user_id=100,
|
||||||
|
content="Hello world.",
|
||||||
|
embed_texts=[],
|
||||||
|
image_urls=[],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
await uningest_message(db, ingest_cache, 602)
|
||||||
|
|
||||||
|
assert ingest_cache.pop(602) is None
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Discord-boundary integration tests."""
|
||||||
@@ -0,0 +1,133 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Fixtures used only by Simcord-backed Discord integration tests."""
|
||||||
|
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero.bot import Crabstero
|
||||||
|
from crabstero.database import Database
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import AsyncGenerator
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from simcord import ChannelHandle, Env, GuildHandle, MemberActor
|
||||||
|
|
||||||
|
type StartWordsForChannel = Callable[[Database, int], Awaitable[list[str]]]
|
||||||
|
type MakeSimcordTextChannel = Callable[["Env"], Awaitable["SimcordTextChannel"]]
|
||||||
|
type MakeSimcordMemberChannel = Callable[["Env"], Awaitable["SimcordMemberChannel"]]
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SimcordTextChannel:
|
||||||
|
"""Guild and text channel created inside a running Simcord environment."""
|
||||||
|
|
||||||
|
guild: GuildHandle
|
||||||
|
channel: ChannelHandle
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class SimcordMemberChannel:
|
||||||
|
"""Guild, text channel, and human member for Discord flow tests."""
|
||||||
|
|
||||||
|
guild: GuildHandle
|
||||||
|
member: MemberActor
|
||||||
|
channel: ChannelHandle
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
async def crabstero_bot(tmp_path: Path) -> AsyncGenerator[Crabstero]:
|
||||||
|
"""Yield the Crabstero bot instance inspected by Discord integration tests."""
|
||||||
|
bot = Crabstero(str(tmp_path / "crabstero.db"))
|
||||||
|
try:
|
||||||
|
yield bot
|
||||||
|
finally:
|
||||||
|
bot.ws = None # type: ignore[assignment]
|
||||||
|
await bot.close()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def simcord_bot(crabstero_bot: Crabstero) -> Crabstero:
|
||||||
|
"""Expose Crabstero under the fixture name required by Simcord."""
|
||||||
|
return crabstero_bot
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def make_simcord_text_channel() -> MakeSimcordTextChannel:
|
||||||
|
"""Return a factory for guild text channels in any Simcord environment."""
|
||||||
|
|
||||||
|
async def make(env: Env) -> SimcordTextChannel:
|
||||||
|
guild = env.create_guild()
|
||||||
|
await env.settle()
|
||||||
|
channel = guild.create_text_channel("general")
|
||||||
|
await env.settle()
|
||||||
|
return SimcordTextChannel(guild, channel)
|
||||||
|
|
||||||
|
return make
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
async def simcord_text_channel(
|
||||||
|
simcord_env: Env,
|
||||||
|
make_simcord_text_channel: MakeSimcordTextChannel,
|
||||||
|
) -> SimcordTextChannel:
|
||||||
|
"""Create one guild text channel in the default Simcord environment."""
|
||||||
|
return await make_simcord_text_channel(simcord_env)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def make_simcord_member_channel() -> MakeSimcordMemberChannel:
|
||||||
|
"""Return a factory for guild/member/channel triples in any Simcord env."""
|
||||||
|
|
||||||
|
async def make(env: Env) -> SimcordMemberChannel:
|
||||||
|
guild = env.create_guild()
|
||||||
|
await env.settle()
|
||||||
|
channel = guild.create_text_channel("general")
|
||||||
|
member = guild.add_member(env.create_user("Ada"))
|
||||||
|
await env.settle()
|
||||||
|
return SimcordMemberChannel(guild, member, channel)
|
||||||
|
|
||||||
|
return make
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
async def simcord_member_channel(
|
||||||
|
simcord_env: Env,
|
||||||
|
make_simcord_member_channel: MakeSimcordMemberChannel,
|
||||||
|
) -> SimcordMemberChannel:
|
||||||
|
"""Create one guild text channel and human member in the default Simcord env."""
|
||||||
|
return await make_simcord_member_channel(simcord_env)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def start_words_for_channel() -> StartWordsForChannel:
|
||||||
|
"""Return persisted Markov start-word occurrences for one Discord channel."""
|
||||||
|
|
||||||
|
async def read(db: Database, channel_id: int) -> list[str]:
|
||||||
|
async with db._connection.execute(
|
||||||
|
"SELECT word, occurrences"
|
||||||
|
" FROM markov_start_words"
|
||||||
|
" WHERE channel_id = ?"
|
||||||
|
" ORDER BY word",
|
||||||
|
(channel_id,),
|
||||||
|
) as cursor:
|
||||||
|
rows = await cursor.fetchall()
|
||||||
|
return [str(row[0]) for row in rows for _ in range(int(row[1]))]
|
||||||
|
|
||||||
|
return read
|
||||||
@@ -0,0 +1,365 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Simcord integration tests for Crabstero bot lifecycle behavior."""
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING, Any, cast
|
||||||
|
|
||||||
|
import discord
|
||||||
|
import pytest
|
||||||
|
from discord import app_commands
|
||||||
|
from discord.ext import commands
|
||||||
|
from simcord import run
|
||||||
|
|
||||||
|
from crabstero import metrics
|
||||||
|
from crabstero.bot import Crabstero, TrackedModal, TrackedView
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from simcord import Env
|
||||||
|
|
||||||
|
from tests.integration.discord.conftest import (
|
||||||
|
MakeSimcordMemberChannel,
|
||||||
|
SimcordMemberChannel,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _counter_value(counter: Any, **labels: str) -> float:
|
||||||
|
"""Return the current value for a labelled Prometheus counter."""
|
||||||
|
return float(counter.labels(**labels)._value.get())
|
||||||
|
|
||||||
|
|
||||||
|
class _LifecycleResource:
|
||||||
|
"""Resource double that records lifecycle cleanup calls."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
name: str,
|
||||||
|
events: list[str],
|
||||||
|
*,
|
||||||
|
fail_once: bool = False,
|
||||||
|
) -> None:
|
||||||
|
"""Store the event name and shared event log."""
|
||||||
|
self._name = name
|
||||||
|
self._events = events
|
||||||
|
self._fail_once = fail_once
|
||||||
|
|
||||||
|
async def _record(self) -> None:
|
||||||
|
"""Record a cleanup call and optionally fail once."""
|
||||||
|
self._events.append(self._name)
|
||||||
|
if self._fail_once:
|
||||||
|
self._fail_once = False
|
||||||
|
raise RuntimeError(f"{self._name} cleanup failed")
|
||||||
|
|
||||||
|
async def close(self) -> None:
|
||||||
|
"""Record a close call."""
|
||||||
|
await self._record()
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
"""Record a stop call."""
|
||||||
|
await self._record()
|
||||||
|
|
||||||
|
|
||||||
|
class TestSetupHook:
|
||||||
|
"""Bot startup wires Discord cogs and slash commands under Simcord."""
|
||||||
|
|
||||||
|
def test_db_before_setup_raises(self, tmp_path: Path) -> None:
|
||||||
|
"""The database property is unavailable before setup_hook runs."""
|
||||||
|
bot = Crabstero(str(tmp_path / "not-started.db"))
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="Database is not initialized"):
|
||||||
|
_ = bot.db
|
||||||
|
|
||||||
|
async def test_loads_cogs_and_syncs_commands(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""setup_hook loads expected cogs and syncs app commands."""
|
||||||
|
assert set(crabstero_bot.cogs) == {
|
||||||
|
"InteractionCog",
|
||||||
|
"MessageCog",
|
||||||
|
"ServerEventsCog",
|
||||||
|
}
|
||||||
|
|
||||||
|
commands = simcord_env.backend.commands[None]
|
||||||
|
assert {name for name, _ in commands} == {"forgetme", "pingme"}
|
||||||
|
|
||||||
|
application_id = simcord_env.backend.application_id
|
||||||
|
http_routes = [f"{method} {path}" for method, path, _ in simcord_env.http_log]
|
||||||
|
assert f"GET /applications/{application_id}/commands" in http_routes
|
||||||
|
assert f"PUT /applications/{application_id}/commands" in http_routes
|
||||||
|
|
||||||
|
async def test_metrics_server_lifecycle_starts_and_stops(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
"""A configured metrics server is started and stopped with the bot."""
|
||||||
|
bot = Crabstero(
|
||||||
|
str(tmp_path / "metrics.db"),
|
||||||
|
metrics_address=metrics.TcpMetricsAddress("127.0.0.1", 0),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
async with run(bot):
|
||||||
|
assert bot._metrics_server is not None
|
||||||
|
assert bot._metrics_server.port > 0
|
||||||
|
finally:
|
||||||
|
bot.ws = None # type: ignore[assignment]
|
||||||
|
await bot.close()
|
||||||
|
|
||||||
|
|
||||||
|
class TestClose:
|
||||||
|
"""Bot shutdown ordering."""
|
||||||
|
|
||||||
|
async def test_closes_discord_before_local_resources(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Discord is closed before local resources are released."""
|
||||||
|
events: list[str] = []
|
||||||
|
bot = Crabstero(str(tmp_path / "close-order.db"))
|
||||||
|
original_close = commands.Bot.close
|
||||||
|
|
||||||
|
async def close_discord(self: commands.Bot) -> None:
|
||||||
|
events.append("discord")
|
||||||
|
await original_close(self)
|
||||||
|
|
||||||
|
monkeypatch.setattr(commands.Bot, "close", close_discord)
|
||||||
|
cast("Any", bot).ingest_cache = _LifecycleResource("cache", events)
|
||||||
|
cast("Any", bot)._metrics_server = _LifecycleResource("metrics", events)
|
||||||
|
cast("Any", bot)._db = _LifecycleResource("db", events)
|
||||||
|
|
||||||
|
await bot.close()
|
||||||
|
|
||||||
|
assert events == ["discord", "cache", "metrics", "db"]
|
||||||
|
|
||||||
|
async def test_local_resources_close_when_discord_close_fails(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Local resources are released even if Discord close raises."""
|
||||||
|
events: list[str] = []
|
||||||
|
bot = Crabstero(str(tmp_path / "close-failure.db"))
|
||||||
|
|
||||||
|
async def close_discord(_self: commands.Bot) -> None:
|
||||||
|
events.append("discord")
|
||||||
|
raise RuntimeError("discord close failed")
|
||||||
|
|
||||||
|
monkeypatch.setattr(commands.Bot, "close", close_discord)
|
||||||
|
cast("Any", bot).ingest_cache = _LifecycleResource("cache", events)
|
||||||
|
cast("Any", bot)._metrics_server = _LifecycleResource("metrics", events)
|
||||||
|
cast("Any", bot)._db = _LifecycleResource("db", events)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="discord close failed"):
|
||||||
|
await bot.close()
|
||||||
|
|
||||||
|
assert events == ["discord", "cache", "metrics", "db"]
|
||||||
|
|
||||||
|
async def test_local_resource_cleanup_is_retried_after_local_failure(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Local cleanup can be retried after Discord has already closed."""
|
||||||
|
events: list[str] = []
|
||||||
|
bot = Crabstero(str(tmp_path / "close-retry.db"))
|
||||||
|
original_close = commands.Bot.close
|
||||||
|
|
||||||
|
async def close_discord(self: commands.Bot) -> None:
|
||||||
|
events.append("discord")
|
||||||
|
await original_close(self)
|
||||||
|
|
||||||
|
monkeypatch.setattr(commands.Bot, "close", close_discord)
|
||||||
|
cast("Any", bot).ingest_cache = _LifecycleResource(
|
||||||
|
"cache",
|
||||||
|
events,
|
||||||
|
fail_once=True,
|
||||||
|
)
|
||||||
|
cast("Any", bot)._metrics_server = _LifecycleResource("metrics", events)
|
||||||
|
cast("Any", bot)._db = _LifecycleResource("db", events)
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="cache cleanup failed"):
|
||||||
|
await bot.close()
|
||||||
|
|
||||||
|
assert events == ["discord", "cache", "metrics", "db"]
|
||||||
|
|
||||||
|
await bot.close()
|
||||||
|
|
||||||
|
assert events == ["discord", "cache", "metrics", "db", "cache"]
|
||||||
|
|
||||||
|
|
||||||
|
class TestMetrics:
|
||||||
|
"""Bot-level metrics are incremented by lifecycle error paths."""
|
||||||
|
|
||||||
|
def test_dispatch_increments_discord_event_metric(
|
||||||
|
self,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""Dispatch increments the labelled Discord event counter."""
|
||||||
|
event = "codex_lifecycle_metric"
|
||||||
|
before = _counter_value(metrics.DISCORD_EVENTS, event=event)
|
||||||
|
|
||||||
|
crabstero_bot.dispatch(event)
|
||||||
|
|
||||||
|
assert _counter_value(metrics.DISCORD_EVENTS, event=event) == before + 1
|
||||||
|
|
||||||
|
async def test_on_error_increments_event_error_metric(
|
||||||
|
self,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""Unhandled event listener failures are counted by event name."""
|
||||||
|
source = "on_message"
|
||||||
|
before = _counter_value(metrics.ERRORS, source=source)
|
||||||
|
|
||||||
|
await crabstero_bot.on_error(source)
|
||||||
|
|
||||||
|
assert _counter_value(metrics.ERRORS, source=source) == before + 1
|
||||||
|
|
||||||
|
async def test_tracked_view_and_modal_errors_increment_metrics(self) -> None:
|
||||||
|
"""Tracked UI error handlers increment their error counters."""
|
||||||
|
view_before = _counter_value(metrics.ERRORS, source="view")
|
||||||
|
modal_before = _counter_value(metrics.ERRORS, source="modal")
|
||||||
|
interaction = cast("discord.Interaction", object())
|
||||||
|
button = cast("discord.ui.Item[TrackedView]", discord.ui.Button(label="Run"))
|
||||||
|
|
||||||
|
await TrackedView().on_error(interaction, RuntimeError("view failed"), button)
|
||||||
|
await TrackedModal(title="Tracked").on_error(
|
||||||
|
interaction,
|
||||||
|
RuntimeError("modal failed"),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert _counter_value(metrics.ERRORS, source="view") == view_before + 1
|
||||||
|
assert _counter_value(metrics.ERRORS, source="modal") == modal_before + 1
|
||||||
|
|
||||||
|
|
||||||
|
class TestCommandErrors:
|
||||||
|
"""Unhandled app command errors produce an ephemeral fallback response."""
|
||||||
|
|
||||||
|
async def test_app_command_error_sends_ephemeral_fallback(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""A failing slash command is captured and answered by tree.on_error."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
|
||||||
|
async def raise_on_set_flag(
|
||||||
|
_entity_type: str,
|
||||||
|
_entity_id: str,
|
||||||
|
_flag_name: str,
|
||||||
|
) -> None:
|
||||||
|
raise RuntimeError("database unavailable")
|
||||||
|
|
||||||
|
monkeypatch.setattr(crabstero_bot.db, "set_flag", raise_on_set_flag)
|
||||||
|
errors_before = _counter_value(metrics.ERRORS, source="command")
|
||||||
|
|
||||||
|
result = await member.slash(channel, "pingme")
|
||||||
|
|
||||||
|
assert result.response is not None
|
||||||
|
assert result.response.ephemeral is True
|
||||||
|
assert result.response.content == (
|
||||||
|
"I encountered an error while processing this command."
|
||||||
|
" Please try again later."
|
||||||
|
)
|
||||||
|
assert _counter_value(metrics.ERRORS, source="command") == errors_before + 1
|
||||||
|
|
||||||
|
async def test_app_command_error_mentions_developer_when_metrics_are_enabled(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
make_simcord_member_channel: MakeSimcordMemberChannel,
|
||||||
|
) -> None:
|
||||||
|
"""Configured metrics switch command errors to the notified-developer copy."""
|
||||||
|
bot = Crabstero(
|
||||||
|
str(tmp_path / "metrics-command-errors.db"),
|
||||||
|
metrics_address=metrics.TcpMetricsAddress("127.0.0.1", 0),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
async with run(bot) as env:
|
||||||
|
context = await make_simcord_member_channel(env)
|
||||||
|
member = context.member
|
||||||
|
channel = context.channel
|
||||||
|
|
||||||
|
async def raise_on_set_flag(
|
||||||
|
_entity_type: str,
|
||||||
|
_entity_id: str,
|
||||||
|
_flag_name: str,
|
||||||
|
) -> None:
|
||||||
|
raise RuntimeError("database unavailable")
|
||||||
|
|
||||||
|
monkeypatch.setattr(bot.db, "set_flag", raise_on_set_flag)
|
||||||
|
errors_before = _counter_value(metrics.ERRORS, source="command")
|
||||||
|
|
||||||
|
result = await member.slash(channel, "pingme")
|
||||||
|
|
||||||
|
assert result.response is not None
|
||||||
|
assert result.response.ephemeral is True
|
||||||
|
assert result.response.content == (
|
||||||
|
"I encountered an error while processing this command."
|
||||||
|
" The developer has been notified,"
|
||||||
|
" please try again later."
|
||||||
|
)
|
||||||
|
assert (
|
||||||
|
_counter_value(metrics.ERRORS, source="command")
|
||||||
|
== errors_before + 1
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
bot.ws = None # type: ignore[assignment]
|
||||||
|
await bot.close()
|
||||||
|
|
||||||
|
async def test_app_command_error_after_defer_sends_ephemeral_followup(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
make_simcord_member_channel: MakeSimcordMemberChannel,
|
||||||
|
) -> None:
|
||||||
|
"""A command that has already acknowledged uses a followup fallback."""
|
||||||
|
bot = Crabstero(str(tmp_path / "deferred-command-errors.db"))
|
||||||
|
|
||||||
|
@app_commands.command(
|
||||||
|
name="deferboom",
|
||||||
|
description="Fail after acknowledging the interaction.",
|
||||||
|
)
|
||||||
|
async def deferboom(interaction: discord.Interaction) -> None:
|
||||||
|
await interaction.response.defer(ephemeral=True)
|
||||||
|
raise RuntimeError("deferred command failed")
|
||||||
|
|
||||||
|
bot.tree.add_command(deferboom)
|
||||||
|
try:
|
||||||
|
async with run(bot) as env:
|
||||||
|
context = await make_simcord_member_channel(env)
|
||||||
|
errors_before = _counter_value(metrics.ERRORS, source="command")
|
||||||
|
|
||||||
|
result = await context.member.slash(context.channel, "deferboom")
|
||||||
|
|
||||||
|
assert result.deferred is True
|
||||||
|
assert result.response is None
|
||||||
|
assert len(result.followups) == 1
|
||||||
|
followup = result.followups[0]
|
||||||
|
assert followup.ephemeral is True
|
||||||
|
assert followup.content == (
|
||||||
|
"I encountered an error while processing this command."
|
||||||
|
" Please try again later."
|
||||||
|
)
|
||||||
|
assert (
|
||||||
|
_counter_value(metrics.ERRORS, source="command")
|
||||||
|
== errors_before + 1
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
bot.ws = None # type: ignore[assignment]
|
||||||
|
await bot.close()
|
||||||
@@ -0,0 +1,276 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Simcord integration tests for Crabstero slash commands."""
|
||||||
|
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero.database import ChannelImage, StartWord, Transition
|
||||||
|
from crabstero.flags import Flag
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from simcord import Env
|
||||||
|
|
||||||
|
from crabstero.bot import Crabstero
|
||||||
|
from crabstero.database import Database
|
||||||
|
from tests.integration.discord.conftest import SimcordMemberChannel
|
||||||
|
|
||||||
|
type SeedForgetmeData = Callable[["Database", int], Awaitable[None]]
|
||||||
|
type ForgetmeRowsForUser = Callable[["Database", int], Awaitable[tuple[int, int, int]]]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def seed_forgetme_data() -> SeedForgetmeData:
|
||||||
|
"""Return a seeder for user-owned and unrelated /forgetme database rows."""
|
||||||
|
|
||||||
|
async def seed(db: Database, user_id: int) -> None:
|
||||||
|
await db.add_markov_data(
|
||||||
|
[
|
||||||
|
StartWord(10, user_id, "delete"),
|
||||||
|
StartWord(10, 999, "keep"),
|
||||||
|
],
|
||||||
|
[
|
||||||
|
Transition(10, user_id, "delete", "me."),
|
||||||
|
Transition(10, 999, "keep", "me."),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
await db.add_images(
|
||||||
|
[
|
||||||
|
ChannelImage(10, user_id, "https://example.com/delete.png"),
|
||||||
|
ChannelImage(10, 999, "https://example.com/keep.png"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
return seed
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def forgetme_rows_for_user() -> ForgetmeRowsForUser:
|
||||||
|
"""Return a reader for Markov/image row counts owned by one user."""
|
||||||
|
|
||||||
|
async def read(db: Database, user_id: int) -> tuple[int, int, int]:
|
||||||
|
async with db._connection.execute(
|
||||||
|
"SELECT COUNT(*) FROM markov_start_words WHERE user_id = ?",
|
||||||
|
(user_id,),
|
||||||
|
) as cursor:
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
assert row is not None
|
||||||
|
start_words = row[0]
|
||||||
|
async with db._connection.execute(
|
||||||
|
"SELECT COUNT(*) FROM markov_transitions WHERE user_id = ?",
|
||||||
|
(user_id,),
|
||||||
|
) as cursor:
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
assert row is not None
|
||||||
|
transitions = row[0]
|
||||||
|
async with db._connection.execute(
|
||||||
|
"SELECT COUNT(*) FROM channel_images WHERE user_id = ?",
|
||||||
|
(user_id,),
|
||||||
|
) as cursor:
|
||||||
|
row = await cursor.fetchone()
|
||||||
|
assert row is not None
|
||||||
|
images = row[0]
|
||||||
|
return int(start_words), int(transitions), int(images)
|
||||||
|
|
||||||
|
return read
|
||||||
|
|
||||||
|
|
||||||
|
class TestPingMe:
|
||||||
|
"""The /pingme command toggles persisted user opt-in state."""
|
||||||
|
|
||||||
|
async def test_first_call_opts_in(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""The first /pingme call sets allowPings and responds ephemerally."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
|
||||||
|
result = await member.slash(channel, "pingme")
|
||||||
|
|
||||||
|
assert result.response is not None
|
||||||
|
assert result.response.ephemeral is True
|
||||||
|
assert "I will now ping you" in result.response.content
|
||||||
|
assert await crabstero_bot.db.is_flag_set(
|
||||||
|
"user",
|
||||||
|
str(member.id),
|
||||||
|
Flag.ALLOW_PINGS,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_second_call_opts_out(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""The second /pingme call clears allowPings and responds ephemerally."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await member.slash(channel, "pingme")
|
||||||
|
|
||||||
|
result = await member.slash(channel, "pingme")
|
||||||
|
|
||||||
|
assert result.response is not None
|
||||||
|
assert result.response.ephemeral is True
|
||||||
|
assert "I will no longer ping you" in result.response.content
|
||||||
|
assert not await crabstero_bot.db.is_flag_set(
|
||||||
|
"user",
|
||||||
|
str(member.id),
|
||||||
|
Flag.ALLOW_PINGS,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestForgetMe:
|
||||||
|
"""The /forgetme command confirms, cancels, and times out hermetically."""
|
||||||
|
|
||||||
|
async def test_initial_response_has_confirmation_buttons(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
) -> None:
|
||||||
|
"""The first /forgetme response is ephemeral and asks for confirmation."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
|
||||||
|
result = await member.slash(channel, "forgetme")
|
||||||
|
|
||||||
|
assert result.response is not None
|
||||||
|
assert result.response.ephemeral is True
|
||||||
|
assert "Would you like to proceed?" in result.response.content
|
||||||
|
labels = [
|
||||||
|
component["label"]
|
||||||
|
for row in result.response.components
|
||||||
|
for component in row["components"]
|
||||||
|
]
|
||||||
|
assert labels == ["Confirm", "Cancel"]
|
||||||
|
|
||||||
|
async def test_confirm_deletes_user_data_and_sets_no_ingest(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
seed_forgetme_data: SeedForgetmeData,
|
||||||
|
forgetme_rows_for_user: ForgetmeRowsForUser,
|
||||||
|
) -> None:
|
||||||
|
"""Confirming /forgetme deletes user data and persists noIngest."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await seed_forgetme_data(crabstero_bot.db, member.id)
|
||||||
|
await crabstero_bot.db.set_flag("user", str(member.id), Flag.ALLOW_PINGS)
|
||||||
|
await crabstero_bot.db.set_flag("user", str(member.id), Flag.NO_REPLY)
|
||||||
|
prompt = (await member.slash(channel, "forgetme")).response
|
||||||
|
assert prompt is not None
|
||||||
|
|
||||||
|
result = await member.click(prompt, label="Confirm")
|
||||||
|
|
||||||
|
assert result.response is not None
|
||||||
|
assert result.response.ephemeral is True
|
||||||
|
assert "I have deleted your data" in result.response.content
|
||||||
|
assert result.response.components == []
|
||||||
|
assert await forgetme_rows_for_user(crabstero_bot.db, member.id) == (0, 0, 0)
|
||||||
|
assert await forgetme_rows_for_user(crabstero_bot.db, 999) == (1, 1, 1)
|
||||||
|
assert await crabstero_bot.db.is_flag_set(
|
||||||
|
"user",
|
||||||
|
str(member.id),
|
||||||
|
Flag.NO_INGEST,
|
||||||
|
)
|
||||||
|
assert not await crabstero_bot.db.is_flag_set(
|
||||||
|
"user",
|
||||||
|
str(member.id),
|
||||||
|
Flag.ALLOW_PINGS,
|
||||||
|
)
|
||||||
|
assert not await crabstero_bot.db.is_flag_set(
|
||||||
|
"user",
|
||||||
|
str(member.id),
|
||||||
|
Flag.NO_REPLY,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_cancel_leaves_user_data_and_flags(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
seed_forgetme_data: SeedForgetmeData,
|
||||||
|
forgetme_rows_for_user: ForgetmeRowsForUser,
|
||||||
|
) -> None:
|
||||||
|
"""Cancelling /forgetme leaves data and user flags unchanged."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await seed_forgetme_data(crabstero_bot.db, member.id)
|
||||||
|
await crabstero_bot.db.set_flag("user", str(member.id), Flag.ALLOW_PINGS)
|
||||||
|
await crabstero_bot.db.set_flag("user", str(member.id), Flag.NO_REPLY)
|
||||||
|
prompt = (await member.slash(channel, "forgetme")).response
|
||||||
|
assert prompt is not None
|
||||||
|
|
||||||
|
result = await member.click(prompt, label="Cancel")
|
||||||
|
|
||||||
|
assert result.response is not None
|
||||||
|
assert result.response.ephemeral is True
|
||||||
|
assert result.response.content == (
|
||||||
|
"Action cancelled. I have not modified your data."
|
||||||
|
)
|
||||||
|
assert result.response.components == []
|
||||||
|
assert await forgetme_rows_for_user(crabstero_bot.db, member.id) == (1, 1, 1)
|
||||||
|
assert await crabstero_bot.db.is_flag_set(
|
||||||
|
"user",
|
||||||
|
str(member.id),
|
||||||
|
Flag.ALLOW_PINGS,
|
||||||
|
)
|
||||||
|
assert await crabstero_bot.db.is_flag_set(
|
||||||
|
"user",
|
||||||
|
str(member.id),
|
||||||
|
Flag.NO_REPLY,
|
||||||
|
)
|
||||||
|
assert not await crabstero_bot.db.is_flag_set(
|
||||||
|
"user",
|
||||||
|
str(member.id),
|
||||||
|
Flag.NO_INGEST,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_already_forgotten_user_gets_terminal_response(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""A noIngest user does not get another confirmation view."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await crabstero_bot.db.set_flag("user", str(member.id), Flag.NO_INGEST)
|
||||||
|
|
||||||
|
result = await member.slash(channel, "forgetme")
|
||||||
|
|
||||||
|
assert result.response is not None
|
||||||
|
assert result.response.ephemeral is True
|
||||||
|
assert result.response.content == (
|
||||||
|
"I have already removed your data and I am not using your messages."
|
||||||
|
)
|
||||||
|
assert result.response.components == []
|
||||||
|
|
||||||
|
async def test_confirmation_timeout_removes_view(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
) -> None:
|
||||||
|
"""The /forgetme view timeout edits the original response without sleeping."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
prompt = (await member.slash(channel, "forgetme")).response
|
||||||
|
assert prompt is not None
|
||||||
|
|
||||||
|
await simcord_env.advance_time(181)
|
||||||
|
|
||||||
|
assert prompt.content == (
|
||||||
|
"This timed out. Run `/forgetme` again if you still want to."
|
||||||
|
)
|
||||||
|
assert prompt.components == []
|
||||||
@@ -0,0 +1,470 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Simcord integration tests for Discord message events."""
|
||||||
|
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import discord
|
||||||
|
import pytest
|
||||||
|
from simcord import run
|
||||||
|
|
||||||
|
from crabstero.bot import Crabstero
|
||||||
|
from crabstero.database import ChannelImage, StartWord, Transition
|
||||||
|
from crabstero.flags import EntityType, Flag
|
||||||
|
from crabstero.messages import ingest_message, reply_to_message
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from simcord import Env
|
||||||
|
|
||||||
|
from crabstero.database import Database
|
||||||
|
from tests.integration.discord.conftest import (
|
||||||
|
MakeSimcordMemberChannel,
|
||||||
|
SimcordMemberChannel,
|
||||||
|
StartWordsForChannel,
|
||||||
|
)
|
||||||
|
|
||||||
|
type SeedReply = Callable[["Database", int], Awaitable[None]]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def seed_reply() -> SeedReply:
|
||||||
|
"""Return a seeder for deterministic generated replies in one channel."""
|
||||||
|
|
||||||
|
async def seed(db: Database, channel_id: int) -> None:
|
||||||
|
await db.add_markov_data(
|
||||||
|
[StartWord(channel_id, 123, "Generated")],
|
||||||
|
[Transition(channel_id, 123, "Generated", "reply.")],
|
||||||
|
)
|
||||||
|
|
||||||
|
return seed
|
||||||
|
|
||||||
|
|
||||||
|
class TestMessageIngestion:
|
||||||
|
"""Normal Discord message flow populates the real database."""
|
||||||
|
|
||||||
|
async def test_normal_guild_user_message_is_ingested(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""A default guild text message adds Markov data."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
|
||||||
|
await member.send(channel, "Alpha beta.")
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == ["Alpha"]
|
||||||
|
assert (
|
||||||
|
await crabstero_bot.db.get_random_next_word(channel.id, "Alpha") == "beta."
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_bot_messages_are_ignored(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""A bot-authored gateway message is not ingested."""
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
|
||||||
|
simcord_env.backend.create_message(
|
||||||
|
channel.id,
|
||||||
|
simcord_env.backend.bot_user.id,
|
||||||
|
"Ignore bot.",
|
||||||
|
)
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
|
||||||
|
|
||||||
|
async def test_dm_messages_are_ignored(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""DM messages are outside the guild channel types Crabstero handles."""
|
||||||
|
user = simcord_env.create_user("Ada")
|
||||||
|
|
||||||
|
await user.send_dm("Direct message.")
|
||||||
|
|
||||||
|
assert (
|
||||||
|
await start_words_for_channel(
|
||||||
|
crabstero_bot.db,
|
||||||
|
user.dm_channel.id,
|
||||||
|
)
|
||||||
|
== []
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_thread_messages_do_not_create_separate_chain(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""Messages inside threads are not ingested under the thread channel id."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
simcord_env.backend.create_thread(channel.id, "thread", member.id)
|
||||||
|
await simcord_env.settle()
|
||||||
|
thread = channel.threads[0]
|
||||||
|
|
||||||
|
await member.send(thread, "Thread only.")
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, thread.id) == []
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
|
||||||
|
|
||||||
|
async def test_embed_message_text_and_image_are_ingested(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""Embed titles, descriptions, and image URLs are ingested."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
|
||||||
|
simcord_env.backend.create_message(
|
||||||
|
channel.id,
|
||||||
|
member.id,
|
||||||
|
embeds=[
|
||||||
|
{
|
||||||
|
"title": "Title words.",
|
||||||
|
"description": "Description words.",
|
||||||
|
"image": {"url": "https://example.com/embed.png"},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == [
|
||||||
|
"Description",
|
||||||
|
"Title",
|
||||||
|
]
|
||||||
|
assert (
|
||||||
|
await crabstero_bot.db.get_random_image(channel.id)
|
||||||
|
== "https://example.com/embed.png"
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_empty_messages_are_not_ingested(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""A message with no content or embeds is ignored."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
|
||||||
|
await member.send(channel)
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
|
||||||
|
|
||||||
|
async def test_ingest_only_mode_ingests_without_replying(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
make_simcord_member_channel: MakeSimcordMemberChannel,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""An ingest-only bot still ingests eligible messages and never replies."""
|
||||||
|
bot = Crabstero(str(tmp_path / "ingest-only.db"), ingest_only=True)
|
||||||
|
try:
|
||||||
|
async with run(bot) as env:
|
||||||
|
context = await make_simcord_member_channel(env)
|
||||||
|
member = context.member
|
||||||
|
channel = context.channel
|
||||||
|
|
||||||
|
await member.send(channel, "Ingest only.")
|
||||||
|
|
||||||
|
assert await start_words_for_channel(bot.db, channel.id) == ["Ingest"]
|
||||||
|
assert [message.content for message in channel.history()] == [
|
||||||
|
"Ingest only.",
|
||||||
|
]
|
||||||
|
finally:
|
||||||
|
bot.ws = None # type: ignore[assignment]
|
||||||
|
await bot.close()
|
||||||
|
|
||||||
|
|
||||||
|
class TestReplies:
|
||||||
|
"""Mentions produce replies through the real Discord message path."""
|
||||||
|
|
||||||
|
async def test_mentioning_bot_sends_seeded_reply(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
seed_reply: SeedReply,
|
||||||
|
) -> None:
|
||||||
|
"""A mention causes a deterministic Markov reply to be posted."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await seed_reply(crabstero_bot.db, channel.id)
|
||||||
|
|
||||||
|
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
|
||||||
|
|
||||||
|
history = channel.history()
|
||||||
|
assert [message.content for message in history] == [
|
||||||
|
f"<@{simcord_env.backend.bot_user.id}> please reply",
|
||||||
|
"Generated reply.",
|
||||||
|
]
|
||||||
|
assert history[-1].reference is not None
|
||||||
|
assert history[-1].author == simcord_env.bot.user
|
||||||
|
|
||||||
|
async def test_mention_reply_can_include_embed_and_image(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
seed_reply: SeedReply,
|
||||||
|
) -> None:
|
||||||
|
"""The optional reply embed path uses generated text and stored images."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await seed_reply(crabstero_bot.db, channel.id)
|
||||||
|
await crabstero_bot.db.add_images(
|
||||||
|
[ChannelImage(channel.id, 123, "https://example.com/reply.png")],
|
||||||
|
)
|
||||||
|
monkeypatch.setattr("crabstero.messages.secrets.randbelow", lambda _upper: 95)
|
||||||
|
|
||||||
|
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
|
||||||
|
|
||||||
|
reply = channel.history()[-1]
|
||||||
|
assert reply.content == "Generated reply."
|
||||||
|
assert len(reply.embeds) == 1
|
||||||
|
assert reply.embeds[0].title == "Generated reply."
|
||||||
|
assert reply.embeds[0].description == "Generated reply."
|
||||||
|
assert reply.embeds[0].image.url == "https://example.com/reply.png"
|
||||||
|
|
||||||
|
async def test_mention_in_thread_uses_parent_channel_chain(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
seed_reply: SeedReply,
|
||||||
|
) -> None:
|
||||||
|
"""A thread mention generates from the parent channel's Markov chain."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await seed_reply(crabstero_bot.db, channel.id)
|
||||||
|
simcord_env.backend.create_thread(channel.id, "thread", member.id)
|
||||||
|
await simcord_env.settle()
|
||||||
|
thread = channel.threads[0]
|
||||||
|
|
||||||
|
await member.send(thread, f"<@{simcord_env.backend.bot_user.id}> thread reply")
|
||||||
|
|
||||||
|
assert [message.content for message in thread.history()] == [
|
||||||
|
f"<@{simcord_env.backend.bot_user.id}> thread reply",
|
||||||
|
"Generated reply.",
|
||||||
|
]
|
||||||
|
|
||||||
|
async def test_allow_pings_flag_allows_generated_user_mentions(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""Generated mentions are allowed for users who opted in."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
target = simcord_env.create_user("Mentioned")
|
||||||
|
await crabstero_bot.db.add_markov_data(
|
||||||
|
[StartWord(channel.id, 123, "Hello")],
|
||||||
|
[Transition(channel.id, 123, "Hello", f"{target.mention}.")],
|
||||||
|
)
|
||||||
|
await crabstero_bot.db.set_flag(
|
||||||
|
EntityType.USER,
|
||||||
|
str(target.id),
|
||||||
|
Flag.ALLOW_PINGS,
|
||||||
|
)
|
||||||
|
|
||||||
|
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
|
||||||
|
|
||||||
|
assert channel.history()[-1].content == f"Hello {target.mention}."
|
||||||
|
|
||||||
|
async def test_no_reply_flag_suppresses_reply(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
seed_reply: SeedReply,
|
||||||
|
) -> None:
|
||||||
|
"""The noReply flag prevents a mention response."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await seed_reply(crabstero_bot.db, channel.id)
|
||||||
|
await crabstero_bot.db.set_flag(
|
||||||
|
EntityType.CHANNEL,
|
||||||
|
str(channel.id),
|
||||||
|
Flag.NO_REPLY,
|
||||||
|
)
|
||||||
|
|
||||||
|
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
|
||||||
|
|
||||||
|
assert [message.content for message in channel.history()] == [
|
||||||
|
f"<@{simcord_env.backend.bot_user.id}> please reply",
|
||||||
|
]
|
||||||
|
|
||||||
|
async def test_missing_send_permission_suppresses_reply(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
seed_reply: SeedReply,
|
||||||
|
) -> None:
|
||||||
|
"""A bot without send_messages permission does not reply."""
|
||||||
|
guild = simcord_env.create_guild()
|
||||||
|
bot_role = guild.roles[simcord_env.backend.bot_user.name]
|
||||||
|
channel = guild.create_text_channel(
|
||||||
|
"readonly",
|
||||||
|
overwrites={
|
||||||
|
bot_role: discord.PermissionOverwrite(send_messages=False),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
member = guild.add_member(simcord_env.create_user("Ada"))
|
||||||
|
await simcord_env.settle()
|
||||||
|
await seed_reply(crabstero_bot.db, channel.id)
|
||||||
|
|
||||||
|
await member.send(channel, f"<@{simcord_env.backend.bot_user.id}> please reply")
|
||||||
|
|
||||||
|
assert [message.content for message in channel.history()] == [
|
||||||
|
f"<@{simcord_env.backend.bot_user.id}> please reply",
|
||||||
|
]
|
||||||
|
|
||||||
|
async def test_direct_dm_reply_is_ignored(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""The reply helper ignores messages outside guilds."""
|
||||||
|
user = simcord_env.create_user("Ada")
|
||||||
|
message = await user.send_dm(f"<@{simcord_env.backend.bot_user.id}> hi")
|
||||||
|
|
||||||
|
await reply_to_message(crabstero_bot.db, message)
|
||||||
|
|
||||||
|
assert [message.content for message in user.dm_channel.history()] == [
|
||||||
|
f"<@{simcord_env.backend.bot_user.id}> hi",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestDeletes:
|
||||||
|
"""Raw delete events reverse recent message ingestion."""
|
||||||
|
|
||||||
|
async def test_delete_reverses_recent_ingest(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""Deleting a cached message removes its Markov rows."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
message = await member.send(channel, "Delete me.")
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == [
|
||||||
|
"Delete",
|
||||||
|
]
|
||||||
|
|
||||||
|
await member.delete(message)
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
|
||||||
|
|
||||||
|
async def test_bulk_delete_reverses_recent_ingests(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""Bulk-deleting cached messages removes their Markov rows."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
first = await member.send(channel, "Bulk one.")
|
||||||
|
second = await member.send(channel, "Bulk two.")
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == [
|
||||||
|
"Bulk",
|
||||||
|
"Bulk",
|
||||||
|
]
|
||||||
|
|
||||||
|
simcord_env.backend.bulk_delete_messages(channel.id, [first.id, second.id])
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
|
||||||
|
|
||||||
|
|
||||||
|
class TestFlagSuppression:
|
||||||
|
"""Message behavior respects persisted noIngest and noReply flags."""
|
||||||
|
|
||||||
|
async def test_no_ingest_user_flag_suppresses_ingestion(
|
||||||
|
self,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""The noIngest flag prevents storing a user's message."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await crabstero_bot.db.set_flag(
|
||||||
|
EntityType.USER,
|
||||||
|
str(member.id),
|
||||||
|
Flag.NO_INGEST,
|
||||||
|
)
|
||||||
|
|
||||||
|
await member.send(channel, "Do not learn.")
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
|
||||||
|
|
||||||
|
async def test_direct_dm_ingest_is_ignored(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""The ingest helper ignores messages outside guilds."""
|
||||||
|
user = simcord_env.create_user("Ada")
|
||||||
|
message = await user.send_dm("Direct helper call.")
|
||||||
|
|
||||||
|
await ingest_message(crabstero_bot.db, message)
|
||||||
|
|
||||||
|
assert (
|
||||||
|
await start_words_for_channel(
|
||||||
|
crabstero_bot.db,
|
||||||
|
user.dm_channel.id,
|
||||||
|
)
|
||||||
|
== []
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_direct_bot_mention_ingest_is_ignored(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""The ingest helper ignores messages that mention the bot."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
message = await member.send(
|
||||||
|
channel,
|
||||||
|
f"<@{simcord_env.backend.bot_user.id}> do not learn",
|
||||||
|
)
|
||||||
|
|
||||||
|
await ingest_message(crabstero_bot.db, message)
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
|
||||||
@@ -0,0 +1,277 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Simcord integration tests for server event ingestion triggers."""
|
||||||
|
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
import discord
|
||||||
|
import pytest
|
||||||
|
from simcord.backend.models import Overwrite
|
||||||
|
from simcord.enums import OverwriteType
|
||||||
|
|
||||||
|
from crabstero.listeners.server_events import ServerEventsCog
|
||||||
|
from crabstero.tasks.ingestion import ingest_channel
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from simcord import Env
|
||||||
|
|
||||||
|
from crabstero.bot import Crabstero
|
||||||
|
from tests.integration.discord.conftest import (
|
||||||
|
SimcordMemberChannel,
|
||||||
|
SimcordTextChannel,
|
||||||
|
StartWordsForChannel,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def server_events_cog(crabstero_bot: Crabstero) -> ServerEventsCog:
|
||||||
|
"""Return the loaded server-events cog from the Simcord-backed bot."""
|
||||||
|
cog = crabstero_bot.get_cog("ServerEventsCog")
|
||||||
|
assert isinstance(cog, ServerEventsCog)
|
||||||
|
return cog
|
||||||
|
|
||||||
|
|
||||||
|
class TestGuildEvents:
|
||||||
|
"""Guild availability and joins queue local channel history ingestion."""
|
||||||
|
|
||||||
|
async def test_guild_available_queues_text_and_voice_channels(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
server_events_cog: ServerEventsCog,
|
||||||
|
) -> None:
|
||||||
|
"""Available guilds enqueue all textable channels."""
|
||||||
|
guild = simcord_env.create_guild()
|
||||||
|
await simcord_env.settle()
|
||||||
|
text = guild.create_text_channel("general")
|
||||||
|
voice = guild.create_voice_channel("voice")
|
||||||
|
await simcord_env.settle()
|
||||||
|
cached_guild = simcord_env.bot.get_guild(guild.id)
|
||||||
|
assert cached_guild is not None
|
||||||
|
|
||||||
|
await server_events_cog.on_guild_available(cached_guild)
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
assert await server_events_cog.bot.db.is_channel_ingested(text.id)
|
||||||
|
assert await server_events_cog.bot.db.is_channel_ingested(voice.id)
|
||||||
|
|
||||||
|
async def test_guild_join_queues_ingestion_and_attempts_owner_notification(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_text_channel: SimcordTextChannel,
|
||||||
|
server_events_cog: ServerEventsCog,
|
||||||
|
) -> None:
|
||||||
|
"""Joining a guild enqueues ingestion and uses the fake owner DM path."""
|
||||||
|
guild = simcord_text_channel.guild
|
||||||
|
channel = simcord_text_channel.channel
|
||||||
|
cached_guild = simcord_env.bot.get_guild(guild.id)
|
||||||
|
assert cached_guild is not None
|
||||||
|
|
||||||
|
await server_events_cog.on_guild_join(cached_guild)
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
assert await server_events_cog.bot.db.is_channel_ingested(channel.id)
|
||||||
|
http_routes = [f"{method} {path}" for method, path, _ in simcord_env.http_log]
|
||||||
|
assert "GET /oauth2/applications/@me" in http_routes
|
||||||
|
assert "POST /users/@me/channels" in http_routes
|
||||||
|
assert any(route.endswith("/messages") for route in http_routes)
|
||||||
|
|
||||||
|
|
||||||
|
class TestPermissionUpdateEvents:
|
||||||
|
"""Permission-changing events trigger ingestion only when relevant."""
|
||||||
|
|
||||||
|
async def test_role_update_queues_only_when_permissions_change_for_bot_role(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_text_channel: SimcordTextChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""Role updates require changed permissions and bot membership."""
|
||||||
|
guild = simcord_text_channel.guild
|
||||||
|
channel = simcord_text_channel.channel
|
||||||
|
role = guild.create_role("reader", permissions=discord.Permissions.none())
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
simcord_env.backend.edit_role(
|
||||||
|
guild.id,
|
||||||
|
role.id,
|
||||||
|
{"permissions": discord.Permissions(view_channel=True).value},
|
||||||
|
)
|
||||||
|
await simcord_env.settle()
|
||||||
|
assert not await crabstero_bot.db.is_channel_ingested(channel.id)
|
||||||
|
|
||||||
|
bot_role = guild.roles[simcord_env.backend.bot_user.name]
|
||||||
|
simcord_env.backend.edit_role(
|
||||||
|
guild.id,
|
||||||
|
bot_role.id,
|
||||||
|
{
|
||||||
|
"permissions": discord.Permissions(
|
||||||
|
view_channel=True,
|
||||||
|
read_message_history=True,
|
||||||
|
).value,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
assert await crabstero_bot.db.is_channel_ingested(channel.id)
|
||||||
|
|
||||||
|
async def test_role_update_same_permissions_does_not_queue(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_text_channel: SimcordTextChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""A role update without permission changes is ignored."""
|
||||||
|
guild = simcord_text_channel.guild
|
||||||
|
channel = simcord_text_channel.channel
|
||||||
|
role = guild.create_role(
|
||||||
|
"reader",
|
||||||
|
permissions=discord.Permissions(read_message_history=True),
|
||||||
|
)
|
||||||
|
simcord_env.backend.add_member_role(
|
||||||
|
guild.id,
|
||||||
|
simcord_env.backend.bot_user.id,
|
||||||
|
role.id,
|
||||||
|
)
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
simcord_env.backend.edit_role(
|
||||||
|
guild.id,
|
||||||
|
role.id,
|
||||||
|
{"permissions": discord.Permissions(read_message_history=True).value},
|
||||||
|
)
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
assert not await crabstero_bot.db.is_channel_ingested(channel.id)
|
||||||
|
|
||||||
|
async def test_channel_update_queues_only_when_overwrites_change(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_text_channel: SimcordTextChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""Channel updates without overwrite changes do not enqueue ingestion."""
|
||||||
|
guild = simcord_text_channel.guild
|
||||||
|
channel = simcord_text_channel.channel
|
||||||
|
|
||||||
|
simcord_env.backend.edit_channel(channel.id, {"topic": "no permission change"})
|
||||||
|
await simcord_env.settle()
|
||||||
|
assert not await crabstero_bot.db.is_channel_ingested(channel.id)
|
||||||
|
|
||||||
|
simcord_env.backend.set_overwrite(
|
||||||
|
channel.id,
|
||||||
|
Overwrite(
|
||||||
|
target_id=guild.default_role.id,
|
||||||
|
type=OverwriteType.ROLE,
|
||||||
|
allow=discord.Permissions(read_message_history=True).value,
|
||||||
|
deny=0,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
assert await crabstero_bot.db.is_channel_ingested(channel.id)
|
||||||
|
|
||||||
|
|
||||||
|
class TestIngestionTasks:
|
||||||
|
"""Server-triggered ingestion task behavior stays local and deterministic."""
|
||||||
|
|
||||||
|
async def test_duplicate_channel_queue_requests_share_one_task(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_text_channel: SimcordTextChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""Queuing the same channel twice before the loop runs creates one task."""
|
||||||
|
channel = simcord_text_channel.channel
|
||||||
|
cached_channel = simcord_env.bot.get_channel(channel.id)
|
||||||
|
assert isinstance(cached_channel, discord.TextChannel)
|
||||||
|
|
||||||
|
crabstero_bot.queue_channel_for_ingestion(cached_channel)
|
||||||
|
crabstero_bot.queue_channel_for_ingestion(cached_channel)
|
||||||
|
|
||||||
|
assert list(crabstero_bot._ingestion_tasks) == [channel.id]
|
||||||
|
await simcord_env.settle()
|
||||||
|
|
||||||
|
async def test_channel_history_ingestion_requires_read_history_permission(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
) -> None:
|
||||||
|
"""A channel missing read history permission is skipped."""
|
||||||
|
guild = simcord_env.create_guild()
|
||||||
|
channel = guild.create_text_channel(
|
||||||
|
"hidden-history",
|
||||||
|
overwrites={
|
||||||
|
guild.default_role: discord.PermissionOverwrite(
|
||||||
|
read_message_history=False,
|
||||||
|
),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
await simcord_env.settle()
|
||||||
|
cached_channel = simcord_env.bot.get_channel(channel.id)
|
||||||
|
assert isinstance(cached_channel, discord.TextChannel)
|
||||||
|
|
||||||
|
await ingest_channel(cached_channel, crabstero_bot.db)
|
||||||
|
|
||||||
|
assert not await crabstero_bot.db.is_channel_ingested(channel.id)
|
||||||
|
|
||||||
|
async def test_channel_history_ingestion_reads_existing_messages(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""Bulk channel ingestion reads historical messages through Discord."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
cached_channel = simcord_env.bot.get_channel(channel.id)
|
||||||
|
assert isinstance(cached_channel, discord.TextChannel)
|
||||||
|
simcord_env.backend.create_message(
|
||||||
|
channel.id,
|
||||||
|
member.id,
|
||||||
|
"Historical message.",
|
||||||
|
broadcast=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
await ingest_channel(cached_channel, crabstero_bot.db)
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == [
|
||||||
|
"Historical",
|
||||||
|
]
|
||||||
|
assert await crabstero_bot.db.is_channel_ingested(channel.id)
|
||||||
|
|
||||||
|
async def test_already_ingested_channel_is_skipped(
|
||||||
|
self,
|
||||||
|
simcord_env: Env,
|
||||||
|
simcord_member_channel: SimcordMemberChannel,
|
||||||
|
crabstero_bot: Crabstero,
|
||||||
|
start_words_for_channel: StartWordsForChannel,
|
||||||
|
) -> None:
|
||||||
|
"""A channel marked ingested is not read again."""
|
||||||
|
member = simcord_member_channel.member
|
||||||
|
channel = simcord_member_channel.channel
|
||||||
|
await crabstero_bot.db.mark_channel_ingested(channel.id)
|
||||||
|
cached_channel = simcord_env.bot.get_channel(channel.id)
|
||||||
|
assert isinstance(cached_channel, discord.TextChannel)
|
||||||
|
simcord_env.backend.create_message(
|
||||||
|
channel.id,
|
||||||
|
member.id,
|
||||||
|
"Historical message.",
|
||||||
|
broadcast=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
await ingest_channel(cached_channel, crabstero_bot.db)
|
||||||
|
|
||||||
|
assert await start_words_for_channel(crabstero_bot.db, channel.id) == []
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Metrics integration tests for Crabstero."""
|
||||||
@@ -0,0 +1,196 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Integration tests for the Prometheus metrics HTTP server."""
|
||||||
|
|
||||||
|
import errno
|
||||||
|
import os
|
||||||
|
import socket
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from http import HTTPStatus
|
||||||
|
from typing import TYPE_CHECKING, Literal, cast
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero.metrics import MetricsServer, TcpMetricsAddress, UnixMetricsAddress
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import AsyncGenerator
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
type MetricsTransport = Literal["tcp", "unix-socket"]
|
||||||
|
|
||||||
|
requires_unix_socket = pytest.mark.skipif(
|
||||||
|
os.name != "posix" or not hasattr(socket, "AF_UNIX"),
|
||||||
|
reason="Unix-socket metrics transports require Unix-domain socket support",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True)
|
||||||
|
class MetricsEndpoint:
|
||||||
|
"""Client details for one running metrics transport."""
|
||||||
|
|
||||||
|
base_url: str
|
||||||
|
unix_socket_path: str | None = None
|
||||||
|
|
||||||
|
def client_session(self) -> aiohttp.ClientSession:
|
||||||
|
"""Create an aiohttp client session for this metrics transport."""
|
||||||
|
connector = (
|
||||||
|
aiohttp.UnixConnector(path=self.unix_socket_path)
|
||||||
|
if self.unix_socket_path is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
return aiohttp.ClientSession(connector=connector)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(
|
||||||
|
params=[
|
||||||
|
pytest.param("tcp", id="tcp"),
|
||||||
|
pytest.param("unix-socket", marks=requires_unix_socket, id="unix-socket"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def metrics_endpoint(
|
||||||
|
request: pytest.FixtureRequest,
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> AsyncGenerator[MetricsEndpoint]:
|
||||||
|
"""Start a MetricsServer on each supported transport."""
|
||||||
|
transport = cast("MetricsTransport", request.param)
|
||||||
|
match transport:
|
||||||
|
case "tcp":
|
||||||
|
server = MetricsServer(TcpMetricsAddress("127.0.0.1", 0))
|
||||||
|
await server.start()
|
||||||
|
endpoint = MetricsEndpoint(f"http://127.0.0.1:{server.port}")
|
||||||
|
case "unix-socket":
|
||||||
|
socket_path = tmp_path / "metrics.sock"
|
||||||
|
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
|
||||||
|
await server.start()
|
||||||
|
endpoint = MetricsEndpoint(
|
||||||
|
"http://crabstero",
|
||||||
|
unix_socket_path=str(socket_path),
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
yield endpoint
|
||||||
|
finally:
|
||||||
|
await server.stop()
|
||||||
|
|
||||||
|
|
||||||
|
class TestMetricsServer:
|
||||||
|
"""HTTP server serves Prometheus metrics on /metrics."""
|
||||||
|
|
||||||
|
@requires_unix_socket
|
||||||
|
def test_unix_socket_server_has_no_tcp_port(self, tmp_path: Path) -> None:
|
||||||
|
"""Unix-socket metrics servers do not expose a TCP port."""
|
||||||
|
server = MetricsServer(UnixMetricsAddress(str(tmp_path / "metrics.sock")))
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="does not have a TCP port"):
|
||||||
|
_ = server.port
|
||||||
|
|
||||||
|
async def test_serves_metrics_endpoint(
|
||||||
|
self,
|
||||||
|
metrics_endpoint: MetricsEndpoint,
|
||||||
|
) -> None:
|
||||||
|
"""GET /metrics returns 200 with metric output containing our metrics."""
|
||||||
|
async with (
|
||||||
|
metrics_endpoint.client_session() as session,
|
||||||
|
session.get(f"{metrics_endpoint.base_url}/metrics") as resp,
|
||||||
|
):
|
||||||
|
assert resp.status == HTTPStatus.OK
|
||||||
|
body = await resp.text()
|
||||||
|
assert "crabstero_build_info" in body
|
||||||
|
|
||||||
|
async def test_non_metrics_path_returns_404(
|
||||||
|
self,
|
||||||
|
metrics_endpoint: MetricsEndpoint,
|
||||||
|
) -> None:
|
||||||
|
"""GET on an unknown path returns 404."""
|
||||||
|
async with (
|
||||||
|
metrics_endpoint.client_session() as session,
|
||||||
|
session.get(f"{metrics_endpoint.base_url}/notfound") as resp,
|
||||||
|
):
|
||||||
|
assert resp.status == HTTPStatus.NOT_FOUND
|
||||||
|
|
||||||
|
@requires_unix_socket
|
||||||
|
async def test_unix_socket_mode_is_applied(self, tmp_path: Path) -> None:
|
||||||
|
"""Configured Unix socket mode is applied after startup."""
|
||||||
|
socket_path = tmp_path / "metrics.sock"
|
||||||
|
server = MetricsServer(UnixMetricsAddress(str(socket_path), mode=0o666))
|
||||||
|
await server.start()
|
||||||
|
try:
|
||||||
|
assert socket_path.stat().st_mode & 0o777 == 0o666
|
||||||
|
finally:
|
||||||
|
await server.stop()
|
||||||
|
|
||||||
|
@requires_unix_socket
|
||||||
|
async def test_stale_unix_socket_path_is_recovered(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
"""Startup removes an abandoned Unix socket file left by a crash."""
|
||||||
|
socket_path = tmp_path / "metrics.sock"
|
||||||
|
with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as stale_socket:
|
||||||
|
stale_socket.bind(str(socket_path))
|
||||||
|
|
||||||
|
server = MetricsServer(UnixMetricsAddress(str(socket_path), mode=0o666))
|
||||||
|
await server.start()
|
||||||
|
try:
|
||||||
|
assert socket_path.stat().st_mode & 0o777 == 0o666
|
||||||
|
connector = aiohttp.UnixConnector(path=str(socket_path))
|
||||||
|
async with (
|
||||||
|
aiohttp.ClientSession(connector=connector) as session,
|
||||||
|
session.get("http://crabstero/metrics") as resp,
|
||||||
|
):
|
||||||
|
assert resp.status == HTTPStatus.OK
|
||||||
|
body = await resp.text()
|
||||||
|
assert "crabstero_build_info" in body
|
||||||
|
finally:
|
||||||
|
await server.stop()
|
||||||
|
|
||||||
|
@requires_unix_socket
|
||||||
|
async def test_active_unix_socket_path_fails(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
"""Startup refuses to replace a Unix socket path that is still in use."""
|
||||||
|
socket_path = tmp_path / "metrics.sock"
|
||||||
|
with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as active_socket:
|
||||||
|
active_socket.bind(str(socket_path))
|
||||||
|
active_socket.listen(1)
|
||||||
|
|
||||||
|
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
|
||||||
|
with pytest.raises(OSError, match="already in use") as exc_info:
|
||||||
|
await server.start()
|
||||||
|
|
||||||
|
assert exc_info.value.errno == errno.EADDRINUSE
|
||||||
|
|
||||||
|
@requires_unix_socket
|
||||||
|
async def test_non_socket_unix_path_fails(self, tmp_path: Path) -> None:
|
||||||
|
"""Startup refuses to replace a non-socket path."""
|
||||||
|
socket_path = tmp_path / "metrics.sock"
|
||||||
|
socket_path.write_text("")
|
||||||
|
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
|
||||||
|
with pytest.raises(FileExistsError):
|
||||||
|
await server.start()
|
||||||
|
|
||||||
|
async def test_starting_server_twice_fails(self) -> None:
|
||||||
|
"""A running metrics server cannot be started twice."""
|
||||||
|
server = MetricsServer(TcpMetricsAddress("127.0.0.1", 0))
|
||||||
|
await server.start()
|
||||||
|
try:
|
||||||
|
with pytest.raises(RuntimeError, match="already running"):
|
||||||
|
await server.start()
|
||||||
|
finally:
|
||||||
|
await server.stop()
|
||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Unit tests for Crabstero."""
|
||||||
@@ -0,0 +1,129 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Unit tests for the IngestCache TTL cache."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero.cache import CachedMessage, IngestCache
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeClock:
|
||||||
|
"""Mutable monotonic clock for deterministic cache expiry tests."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.now = 1000.0
|
||||||
|
|
||||||
|
def __call__(self) -> float:
|
||||||
|
"""Return the current fake monotonic timestamp."""
|
||||||
|
return self.now
|
||||||
|
|
||||||
|
def advance(self, seconds: float) -> None:
|
||||||
|
"""Move the fake clock forward by the given number of seconds."""
|
||||||
|
self.now += seconds
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def cached_message() -> CachedMessage:
|
||||||
|
"""Return a representative cache entry for put/pop behavior tests."""
|
||||||
|
return CachedMessage(
|
||||||
|
channel_id=1,
|
||||||
|
user_id=100,
|
||||||
|
content="Hello world",
|
||||||
|
embed_texts=["embed title"],
|
||||||
|
image_urls=["https://example.com/cat.png"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def fake_clock(monkeypatch: pytest.MonkeyPatch) -> _FakeClock:
|
||||||
|
"""Patch the cache's clock and return a controllable timestamp source."""
|
||||||
|
clock = _FakeClock()
|
||||||
|
monkeypatch.setattr("crabstero.cache.time.monotonic", clock)
|
||||||
|
return clock
|
||||||
|
|
||||||
|
|
||||||
|
class TestIngestCache:
|
||||||
|
"""IngestCache put, pop, and expiry behavior."""
|
||||||
|
|
||||||
|
def test_put_and_pop(self, cached_message: CachedMessage) -> None:
|
||||||
|
"""A cached message can be retrieved by message ID."""
|
||||||
|
cache = IngestCache()
|
||||||
|
cache.put(12345, cached_message)
|
||||||
|
assert cache.pop(12345) is cached_message
|
||||||
|
|
||||||
|
def test_pop_removes_entry(self, cached_message: CachedMessage) -> None:
|
||||||
|
"""Popping an entry removes it from the cache."""
|
||||||
|
cache = IngestCache()
|
||||||
|
cache.put(12345, cached_message)
|
||||||
|
cache.pop(12345)
|
||||||
|
assert cache.pop(12345) is None
|
||||||
|
|
||||||
|
def test_pop_returns_none_for_missing(self) -> None:
|
||||||
|
"""Popping a nonexistent key returns None."""
|
||||||
|
cache = IngestCache()
|
||||||
|
assert cache.pop(99999) is None
|
||||||
|
|
||||||
|
def test_expired_entry_not_returned(
|
||||||
|
self,
|
||||||
|
cached_message: CachedMessage,
|
||||||
|
fake_clock: _FakeClock,
|
||||||
|
) -> None:
|
||||||
|
"""An expired entry is not returned by pop."""
|
||||||
|
cache = IngestCache(ttl_seconds=10)
|
||||||
|
cache.put(12345, cached_message)
|
||||||
|
fake_clock.advance(11)
|
||||||
|
assert cache.pop(12345) is None
|
||||||
|
|
||||||
|
def test_cleanup_removes_expired(
|
||||||
|
self,
|
||||||
|
cached_message: CachedMessage,
|
||||||
|
fake_clock: _FakeClock,
|
||||||
|
) -> None:
|
||||||
|
"""Cleanup evicts all expired entries."""
|
||||||
|
cache = IngestCache(ttl_seconds=10)
|
||||||
|
cache.put(1, cached_message)
|
||||||
|
cache.put(2, cached_message)
|
||||||
|
fake_clock.advance(11)
|
||||||
|
cache._cleanup()
|
||||||
|
assert cache.pop(1) is None
|
||||||
|
assert cache.pop(2) is None
|
||||||
|
|
||||||
|
def test_cleanup_keeps_unexpired(self, cached_message: CachedMessage) -> None:
|
||||||
|
"""Cleanup does not evict entries that are still valid."""
|
||||||
|
cache = IngestCache(ttl_seconds=300)
|
||||||
|
cache.put(1, cached_message)
|
||||||
|
cache._cleanup()
|
||||||
|
assert cache.pop(1) is cached_message
|
||||||
|
|
||||||
|
async def test_start_and_stop(self) -> None:
|
||||||
|
"""The background cleanup task can be started and stopped."""
|
||||||
|
cache = IngestCache()
|
||||||
|
cache.start()
|
||||||
|
assert cache._task is not None
|
||||||
|
|
||||||
|
await cache.stop()
|
||||||
|
|
||||||
|
assert cache._task is None
|
||||||
|
|
||||||
|
async def test_start_is_idempotent(self) -> None:
|
||||||
|
"""Starting an already-started cache keeps the existing cleanup task."""
|
||||||
|
cache = IngestCache()
|
||||||
|
cache.start()
|
||||||
|
try:
|
||||||
|
task = cache._task
|
||||||
|
cache.start()
|
||||||
|
assert cache._task is task
|
||||||
|
finally:
|
||||||
|
await cache.stop()
|
||||||
@@ -0,0 +1,629 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Unit tests for CLI argument parsing, credentials, and systemd integration.
|
||||||
|
|
||||||
|
Tests cover the CLI helpers and runtime orchestration from crabstero.cli.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import os
|
||||||
|
import signal
|
||||||
|
import sys
|
||||||
|
from typing import TYPE_CHECKING, Any, ClassVar, Self
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero import cli
|
||||||
|
from crabstero.cli import (
|
||||||
|
_parse_args,
|
||||||
|
_read_credential,
|
||||||
|
_sd_notify,
|
||||||
|
_watchdog_interval,
|
||||||
|
_watchdog_loop,
|
||||||
|
)
|
||||||
|
from crabstero.metrics import TcpMetricsAddress, UnixMetricsAddress
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import Callable, Coroutine
|
||||||
|
from pathlib import Path
|
||||||
|
from types import TracebackType
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeCrabstero:
|
||||||
|
"""Fake bot that records CLI lifecycle calls without touching Discord.
|
||||||
|
|
||||||
|
Tests can inject callbacks during login/connect and exceptions during
|
||||||
|
login/close to exercise CLI shutdown ordering and propagation paths.
|
||||||
|
"""
|
||||||
|
|
||||||
|
events: ClassVar[list[object]] = []
|
||||||
|
login_error: ClassVar[BaseException | None] = None
|
||||||
|
login_callback: ClassVar[Callable[[], None] | None] = None
|
||||||
|
close_error: ClassVar[BaseException | None] = None
|
||||||
|
connect_callback: ClassVar[Callable[[], None] | None] = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def reset(
|
||||||
|
cls,
|
||||||
|
*,
|
||||||
|
login_error: BaseException | None = None,
|
||||||
|
close_error: BaseException | None = None,
|
||||||
|
) -> list[object]:
|
||||||
|
"""Clear scenario state and return the shared lifecycle event log.
|
||||||
|
|
||||||
|
The event log is shared across fake instances so tests can assert the
|
||||||
|
ordering of construction, login, notifications, connect, and close.
|
||||||
|
"""
|
||||||
|
cls.events = []
|
||||||
|
cls.login_error = login_error
|
||||||
|
cls.login_callback = None
|
||||||
|
cls.close_error = close_error
|
||||||
|
cls.connect_callback = None
|
||||||
|
return cls.events
|
||||||
|
|
||||||
|
def __init__(self, **kwargs: object) -> None:
|
||||||
|
self._closed = False
|
||||||
|
self.events.append(("init", kwargs))
|
||||||
|
|
||||||
|
async def __aenter__(self) -> Self:
|
||||||
|
self.events.append("enter")
|
||||||
|
return self
|
||||||
|
|
||||||
|
async def __aexit__(
|
||||||
|
self,
|
||||||
|
exc_type: type[BaseException] | None,
|
||||||
|
exc: BaseException | None,
|
||||||
|
traceback: TracebackType | None,
|
||||||
|
) -> None:
|
||||||
|
self.events.append("exit")
|
||||||
|
await self.close()
|
||||||
|
|
||||||
|
async def login(self, token: str) -> None:
|
||||||
|
self.events.append(("login", token))
|
||||||
|
callback = type(self).login_callback
|
||||||
|
if callback is not None:
|
||||||
|
callback()
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
if self.login_error is not None:
|
||||||
|
raise self.login_error
|
||||||
|
|
||||||
|
async def connect(self, **kwargs: object) -> None:
|
||||||
|
self.events.append(("connect", kwargs))
|
||||||
|
callback = type(self).connect_callback
|
||||||
|
if callback is not None:
|
||||||
|
callback()
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
async def close(self) -> None:
|
||||||
|
if self._closed:
|
||||||
|
return
|
||||||
|
self.events.append("close")
|
||||||
|
self._closed = True
|
||||||
|
if self.close_error is not None:
|
||||||
|
raise self.close_error
|
||||||
|
|
||||||
|
def is_closed(self) -> bool:
|
||||||
|
return self._closed
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def prepared_main(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> Callable[..., list[object]]:
|
||||||
|
"""Return a factory that prepares cli.main for fake-bot lifecycle tests.
|
||||||
|
|
||||||
|
Each factory call resets the fake bot, routes systemd notifications into
|
||||||
|
the event log, disables watchdog setup, and makes uvloop.run execute the
|
||||||
|
coroutine synchronously through asyncio.run.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def prepare(
|
||||||
|
*,
|
||||||
|
login_error: BaseException | None = None,
|
||||||
|
close_error: BaseException | None = None,
|
||||||
|
) -> list[object]:
|
||||||
|
events = _FakeCrabstero.reset(
|
||||||
|
login_error=login_error,
|
||||||
|
close_error=close_error,
|
||||||
|
)
|
||||||
|
|
||||||
|
def sd_notify(state: str) -> None:
|
||||||
|
events.append(("notify", state))
|
||||||
|
|
||||||
|
monkeypatch.setattr(sys, "argv", ["crabstero", "--token", "cli-token"])
|
||||||
|
monkeypatch.setattr(cli, "Crabstero", _FakeCrabstero)
|
||||||
|
monkeypatch.setattr("crabstero.cli.uvloop.run", asyncio.run)
|
||||||
|
monkeypatch.setattr(cli, "_sd_notify", sd_notify)
|
||||||
|
monkeypatch.setattr(cli, "_watchdog_interval", lambda: None)
|
||||||
|
|
||||||
|
return events
|
||||||
|
|
||||||
|
return prepare
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def signal_handlers(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> dict[int, Callable[[], None]]:
|
||||||
|
"""Patch signal registration and return captured signal callbacks.
|
||||||
|
|
||||||
|
The returned mapping lets tests invoke registered SIGINT/SIGTERM handlers
|
||||||
|
directly and assert that cleanup removes them.
|
||||||
|
"""
|
||||||
|
handlers: dict[int, Callable[[], None]] = {}
|
||||||
|
|
||||||
|
class SignalLoop:
|
||||||
|
"""Minimal event-loop facade that stores installed signal handlers."""
|
||||||
|
|
||||||
|
def add_signal_handler(
|
||||||
|
self,
|
||||||
|
sig: int,
|
||||||
|
callback: Callable[..., None],
|
||||||
|
*args: object,
|
||||||
|
) -> None:
|
||||||
|
handlers[sig] = lambda: callback(*args)
|
||||||
|
|
||||||
|
def remove_signal_handler(self, sig: int) -> None:
|
||||||
|
handlers.pop(sig, None)
|
||||||
|
|
||||||
|
monkeypatch.setattr(asyncio, "get_running_loop", SignalLoop)
|
||||||
|
return handlers
|
||||||
|
|
||||||
|
|
||||||
|
class TestParseArgs:
|
||||||
|
"""Argument parsing with defaults and overrides."""
|
||||||
|
|
||||||
|
def test_token_from_arg(self) -> None:
|
||||||
|
"""--token flag sets the token."""
|
||||||
|
args = _parse_args(["--token", "abc"])
|
||||||
|
assert args.token == "abc" # noqa: S105
|
||||||
|
|
||||||
|
def test_database_default(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
"""Omitting --database-path defaults to 'crabstero.db'."""
|
||||||
|
monkeypatch.delenv("DATABASE_PATH", raising=False)
|
||||||
|
args = _parse_args(["--token", "test"])
|
||||||
|
assert args.database_path == "crabstero.db"
|
||||||
|
|
||||||
|
def test_database_override(self) -> None:
|
||||||
|
"""--database-path overrides the default."""
|
||||||
|
args = _parse_args(["--token", "test", "--database-path", "/custom.db"])
|
||||||
|
assert args.database_path == "/custom.db"
|
||||||
|
|
||||||
|
def test_ingest_only_flag(self) -> None:
|
||||||
|
"""--ingest-only sets ingest_only to True."""
|
||||||
|
args = _parse_args(["--token", "test", "--ingest-only"])
|
||||||
|
assert args.ingest_only is True
|
||||||
|
|
||||||
|
def test_listen_metrics_parses_address(self) -> None:
|
||||||
|
"""--listen-metrics HOST:PORT sets listen_metrics to (host, port)."""
|
||||||
|
args = _parse_args(["--token", "test", "--listen-metrics", "127.0.0.1:9090"])
|
||||||
|
assert args.metrics_address == TcpMetricsAddress("127.0.0.1", 9090)
|
||||||
|
|
||||||
|
def test_listen_metrics_default_none(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
"""Omitting --listen-metrics defaults to None."""
|
||||||
|
monkeypatch.delenv("LISTEN_METRICS", raising=False)
|
||||||
|
args = _parse_args(["--token", "test"])
|
||||||
|
assert args.metrics_address is None
|
||||||
|
|
||||||
|
def test_listen_metrics_ipv6(self) -> None:
|
||||||
|
"""--listen-metrics [::1]:PORT parses IPv6 address correctly."""
|
||||||
|
args = _parse_args(["--token", "test", "--listen-metrics", "[::1]:9090"])
|
||||||
|
assert args.metrics_address == TcpMetricsAddress("[::1]", 9090)
|
||||||
|
|
||||||
|
def test_listen_metrics_from_env(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
"""LISTEN_METRICS env var is used when --listen-metrics is omitted."""
|
||||||
|
monkeypatch.setenv("LISTEN_METRICS", "127.0.0.1:8080")
|
||||||
|
args = _parse_args(["--token", "test"])
|
||||||
|
assert args.metrics_address == TcpMetricsAddress("127.0.0.1", 8080)
|
||||||
|
|
||||||
|
def test_listen_metrics_unix_socket(self) -> None:
|
||||||
|
"""--listen-metrics unix:/path sets listen_metrics to a Unix socket."""
|
||||||
|
args = _parse_args(
|
||||||
|
["--token", "test", "--listen-metrics", "unix:/run/crabstero.sock"],
|
||||||
|
)
|
||||||
|
assert args.metrics_address == UnixMetricsAddress("/run/crabstero.sock")
|
||||||
|
|
||||||
|
def test_listen_metrics_unix_socket_mode(self) -> None:
|
||||||
|
"""--metrics-unix-socket-mode sets the Unix socket file mode."""
|
||||||
|
args = _parse_args(
|
||||||
|
[
|
||||||
|
"--token",
|
||||||
|
"test",
|
||||||
|
"--listen-metrics",
|
||||||
|
"unix:/run/crabstero.sock",
|
||||||
|
"--metrics-unix-socket-mode",
|
||||||
|
"0666",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
assert args.metrics_address == UnixMetricsAddress(
|
||||||
|
"/run/crabstero.sock",
|
||||||
|
mode=0o666,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_listen_metrics_unix_socket_from_env(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""LISTEN_METRICS can enable metrics on a Unix socket."""
|
||||||
|
monkeypatch.setenv("LISTEN_METRICS", "unix:/run/crabstero.sock")
|
||||||
|
args = _parse_args(["--token", "test"])
|
||||||
|
assert args.metrics_address == UnixMetricsAddress("/run/crabstero.sock")
|
||||||
|
|
||||||
|
def test_listen_metrics_unix_socket_mode_from_env(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""METRICS_UNIX_SOCKET_MODE sets the Unix socket file mode."""
|
||||||
|
monkeypatch.setenv("LISTEN_METRICS", "unix:/run/crabstero.sock")
|
||||||
|
monkeypatch.setenv("METRICS_UNIX_SOCKET_MODE", "666")
|
||||||
|
args = _parse_args(["--token", "test"])
|
||||||
|
assert args.metrics_address == UnixMetricsAddress(
|
||||||
|
"/run/crabstero.sock",
|
||||||
|
mode=0o666,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_listen_metrics_invalid_format(self) -> None:
|
||||||
|
"""--listen-metrics with no colon raises SystemExit."""
|
||||||
|
with pytest.raises(SystemExit):
|
||||||
|
_parse_args(["--token", "test", "--listen-metrics", "bad"])
|
||||||
|
|
||||||
|
def test_listen_metrics_invalid_port(self) -> None:
|
||||||
|
"""--listen-metrics with non-integer port raises SystemExit."""
|
||||||
|
with pytest.raises(SystemExit):
|
||||||
|
_parse_args(["--token", "test", "--listen-metrics", "127.0.0.1:abc"])
|
||||||
|
|
||||||
|
def test_listen_metrics_empty_unix_socket_path(self) -> None:
|
||||||
|
"""--listen-metrics unix: with no path raises SystemExit."""
|
||||||
|
with pytest.raises(SystemExit):
|
||||||
|
_parse_args(["--token", "test", "--listen-metrics", "unix:"])
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"mode",
|
||||||
|
[
|
||||||
|
pytest.param("bad", id="not-octal"),
|
||||||
|
pytest.param("0888", id="invalid-octal-digit"),
|
||||||
|
pytest.param("1000", id="too-large"),
|
||||||
|
pytest.param("-1", id="negative"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_metrics_unix_socket_mode_invalid(self, mode: str) -> None:
|
||||||
|
"""Invalid Unix socket modes fail argument parsing."""
|
||||||
|
with pytest.raises(SystemExit):
|
||||||
|
_parse_args(
|
||||||
|
[
|
||||||
|
"--token",
|
||||||
|
"test",
|
||||||
|
"--listen-metrics",
|
||||||
|
"unix:/run/crabstero.sock",
|
||||||
|
"--metrics-unix-socket-mode",
|
||||||
|
mode,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_metrics_unix_socket_mode_requires_unix_socket(self) -> None:
|
||||||
|
"""Unix socket mode cannot be used with a TCP metrics listener."""
|
||||||
|
with pytest.raises(SystemExit):
|
||||||
|
_parse_args(
|
||||||
|
[
|
||||||
|
"--token",
|
||||||
|
"test",
|
||||||
|
"--listen-metrics",
|
||||||
|
"127.0.0.1:9090",
|
||||||
|
"--metrics-unix-socket-mode",
|
||||||
|
"0666",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_metrics_unix_socket_mode_requires_metrics_listener(self) -> None:
|
||||||
|
"""Unix socket mode cannot be used without a metrics listener."""
|
||||||
|
with pytest.raises(SystemExit):
|
||||||
|
_parse_args(
|
||||||
|
["--token", "test", "--metrics-unix-socket-mode", "0666"],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_missing_token_exits(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
"""Missing token causes SystemExit."""
|
||||||
|
monkeypatch.delenv("TOKEN", raising=False)
|
||||||
|
monkeypatch.delenv("CREDENTIALS_DIRECTORY", raising=False)
|
||||||
|
with pytest.raises(SystemExit):
|
||||||
|
_parse_args([])
|
||||||
|
|
||||||
|
|
||||||
|
class TestReadCredential:
|
||||||
|
"""Systemd credential file reading."""
|
||||||
|
|
||||||
|
def test_reads_credential_file(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Reads and strips the credential value from the file."""
|
||||||
|
(tmp_path / "mytoken").write_text(" secret123 \n")
|
||||||
|
monkeypatch.setenv("CREDENTIALS_DIRECTORY", str(tmp_path))
|
||||||
|
assert _read_credential("mytoken") == "secret123"
|
||||||
|
|
||||||
|
def test_returns_none_without_env(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
"""Returns None when CREDENTIALS_DIRECTORY is not set."""
|
||||||
|
monkeypatch.delenv("CREDENTIALS_DIRECTORY", raising=False)
|
||||||
|
assert _read_credential("anything") is None
|
||||||
|
|
||||||
|
def test_returns_none_for_missing_file(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Returns None when the credential file does not exist."""
|
||||||
|
monkeypatch.setenv("CREDENTIALS_DIRECTORY", str(tmp_path))
|
||||||
|
assert _read_credential("nonexistent") is None
|
||||||
|
|
||||||
|
|
||||||
|
class TestSystemdNotify:
|
||||||
|
"""Systemd notification helper behavior."""
|
||||||
|
|
||||||
|
def test_sd_notify_noop_without_socket(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Missing NOTIFY_SOCKET is a no-op."""
|
||||||
|
monkeypatch.delenv("NOTIFY_SOCKET", raising=False)
|
||||||
|
_sd_notify("READY=1")
|
||||||
|
|
||||||
|
|
||||||
|
class TestWatchdog:
|
||||||
|
"""Systemd watchdog helper behavior."""
|
||||||
|
|
||||||
|
def test_watchdog_interval_missing(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Missing WATCHDOG_USEC disables watchdog pings."""
|
||||||
|
monkeypatch.delenv("WATCHDOG_USEC", raising=False)
|
||||||
|
assert _watchdog_interval() is None
|
||||||
|
|
||||||
|
def test_watchdog_interval_valid(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""WATCHDOG_USEC is converted to half the interval in seconds."""
|
||||||
|
monkeypatch.setenv("WATCHDOG_USEC", "4000000")
|
||||||
|
monkeypatch.delenv("WATCHDOG_PID", raising=False)
|
||||||
|
assert _watchdog_interval() == 2
|
||||||
|
|
||||||
|
def test_watchdog_interval_valid_for_matching_pid(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""WATCHDOG_PID enables watchdog pings for the current process."""
|
||||||
|
monkeypatch.setenv("WATCHDOG_USEC", "4000000")
|
||||||
|
monkeypatch.setenv("WATCHDOG_PID", str(os.getpid()))
|
||||||
|
assert _watchdog_interval() == 2
|
||||||
|
|
||||||
|
def test_watchdog_interval_ignores_different_pid(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""WATCHDOG_PID disables pings for non-matching processes."""
|
||||||
|
monkeypatch.setenv("WATCHDOG_USEC", "4000000")
|
||||||
|
monkeypatch.setenv("WATCHDOG_PID", str(os.getpid() + 1))
|
||||||
|
assert _watchdog_interval() is None
|
||||||
|
|
||||||
|
def test_watchdog_interval_ignores_invalid_pid(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Invalid WATCHDOG_PID disables watchdog pings."""
|
||||||
|
monkeypatch.setenv("WATCHDOG_USEC", "4000000")
|
||||||
|
monkeypatch.setenv("WATCHDOG_PID", "invalid")
|
||||||
|
assert _watchdog_interval() is None
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"value",
|
||||||
|
[
|
||||||
|
pytest.param("invalid", id="not-integer"),
|
||||||
|
pytest.param("0", id="zero"),
|
||||||
|
pytest.param("-1", id="negative"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_watchdog_interval_invalid(
|
||||||
|
self,
|
||||||
|
value: str,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""Invalid WATCHDOG_USEC values disable watchdog pings."""
|
||||||
|
monkeypatch.setenv("WATCHDOG_USEC", value)
|
||||||
|
assert _watchdog_interval() is None
|
||||||
|
|
||||||
|
async def test_watchdog_loop_sends_until_stopped(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""_watchdog_loop sends WATCHDOG=1 and exits when stopped."""
|
||||||
|
notifications: list[str] = []
|
||||||
|
stop = asyncio.Event()
|
||||||
|
monkeypatch.setattr(cli, "_sd_notify", notifications.append)
|
||||||
|
|
||||||
|
task = asyncio.create_task(_watchdog_loop(60, stop))
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
stop.set()
|
||||||
|
await asyncio.wait_for(task, timeout=1)
|
||||||
|
|
||||||
|
assert notifications == ["WATCHDOG=1"]
|
||||||
|
|
||||||
|
|
||||||
|
class TestMainLifecycle:
|
||||||
|
"""CLI runtime orchestration."""
|
||||||
|
|
||||||
|
def test_constructs_bot_without_token(
|
||||||
|
self,
|
||||||
|
prepared_main: Callable[..., list[object]],
|
||||||
|
) -> None:
|
||||||
|
"""The CLI keeps the Discord token out of Crabstero construction."""
|
||||||
|
events = prepared_main()
|
||||||
|
|
||||||
|
assert cli.main() == 0
|
||||||
|
|
||||||
|
init_event = events[0]
|
||||||
|
assert isinstance(init_event, tuple)
|
||||||
|
assert init_event == (
|
||||||
|
"init",
|
||||||
|
{
|
||||||
|
"database_path": "crabstero.db",
|
||||||
|
"ingest_only": False,
|
||||||
|
"metrics_address": None,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
init_kwargs = init_event[1]
|
||||||
|
assert isinstance(init_kwargs, dict)
|
||||||
|
assert "token" not in init_kwargs
|
||||||
|
|
||||||
|
def test_main_accepts_explicit_argv(
|
||||||
|
self,
|
||||||
|
prepared_main: Callable[..., list[object]],
|
||||||
|
) -> None:
|
||||||
|
"""main(argv) runs from explicit arguments instead of sys.argv."""
|
||||||
|
events = prepared_main()
|
||||||
|
|
||||||
|
assert cli.main(["--token", "argv-token", "--database-path", "/custom.db"]) == 0
|
||||||
|
|
||||||
|
assert events[0] == (
|
||||||
|
"init",
|
||||||
|
{
|
||||||
|
"database_path": "/custom.db",
|
||||||
|
"ingest_only": False,
|
||||||
|
"metrics_address": None,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert ("login", "argv-token") in events
|
||||||
|
|
||||||
|
def test_ready_sent_after_login_before_connect(
|
||||||
|
self,
|
||||||
|
prepared_main: Callable[..., list[object]],
|
||||||
|
) -> None:
|
||||||
|
"""READY=1 is sent after login succeeds and before connect starts."""
|
||||||
|
events = prepared_main()
|
||||||
|
|
||||||
|
assert cli.main() == 0
|
||||||
|
|
||||||
|
assert events.index(("login", "cli-token")) < events.index(
|
||||||
|
("notify", "READY=1"),
|
||||||
|
)
|
||||||
|
assert events.index(("notify", "READY=1")) < events.index(
|
||||||
|
("connect", {"reconnect": True}),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_stopping_sent_after_clean_connect_return(
|
||||||
|
self,
|
||||||
|
prepared_main: Callable[..., list[object]],
|
||||||
|
) -> None:
|
||||||
|
"""STOPPING=1 is sent when the CLI run exits cleanly after readiness."""
|
||||||
|
events = prepared_main()
|
||||||
|
|
||||||
|
assert cli.main() == 0
|
||||||
|
|
||||||
|
assert events.index(("connect", {"reconnect": True})) < events.index(
|
||||||
|
("notify", "STOPPING=1"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_ready_not_sent_when_login_fails(
|
||||||
|
self,
|
||||||
|
prepared_main: Callable[..., list[object]],
|
||||||
|
) -> None:
|
||||||
|
"""A login failure exits without reporting readiness."""
|
||||||
|
events = prepared_main(login_error=RuntimeError("login failed"))
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="login failed"):
|
||||||
|
cli.main()
|
||||||
|
|
||||||
|
assert ("notify", "READY=1") not in events
|
||||||
|
assert ("notify", "STOPPING=1") not in events
|
||||||
|
|
||||||
|
def test_shutdown_task_exception_is_propagated(
|
||||||
|
self,
|
||||||
|
prepared_main: Callable[..., list[object]],
|
||||||
|
signal_handlers: dict[int, Callable[[], None]],
|
||||||
|
) -> None:
|
||||||
|
"""A completed shutdown task exception is not silently dropped."""
|
||||||
|
events = prepared_main(close_error=RuntimeError("close failed"))
|
||||||
|
|
||||||
|
def request_shutdown() -> None:
|
||||||
|
signal_handlers[signal.SIGTERM]()
|
||||||
|
|
||||||
|
_FakeCrabstero.connect_callback = request_shutdown
|
||||||
|
|
||||||
|
with pytest.raises(RuntimeError, match="close failed"):
|
||||||
|
cli.main()
|
||||||
|
|
||||||
|
assert events.index(("notify", "STOPPING=1")) < events.index("close")
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("sig", "exit_code"),
|
||||||
|
[
|
||||||
|
pytest.param(signal.SIGINT, 130, id="sigint"),
|
||||||
|
pytest.param(signal.SIGTERM, 0, id="sigterm"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_registered_signal_handler_sets_exit_code(
|
||||||
|
self,
|
||||||
|
sig: signal.Signals,
|
||||||
|
exit_code: int,
|
||||||
|
prepared_main: Callable[..., list[object]],
|
||||||
|
signal_handlers: dict[int, Callable[[], None]],
|
||||||
|
) -> None:
|
||||||
|
"""Signal callbacks preserve SIGINT and SIGTERM CLI exit semantics."""
|
||||||
|
events = prepared_main()
|
||||||
|
|
||||||
|
def request_shutdown() -> None:
|
||||||
|
signal_handlers[sig]()
|
||||||
|
|
||||||
|
_FakeCrabstero.connect_callback = request_shutdown
|
||||||
|
|
||||||
|
assert cli.main() == exit_code
|
||||||
|
assert ("notify", "STOPPING=1") in events
|
||||||
|
assert "close" in events
|
||||||
|
assert signal_handlers == {}
|
||||||
|
|
||||||
|
def test_registered_signal_during_login_exits_before_ready(
|
||||||
|
self,
|
||||||
|
prepared_main: Callable[..., list[object]],
|
||||||
|
signal_handlers: dict[int, Callable[[], None]],
|
||||||
|
) -> None:
|
||||||
|
"""A signal during login cancels startup without reporting readiness."""
|
||||||
|
events = prepared_main()
|
||||||
|
|
||||||
|
def request_shutdown() -> None:
|
||||||
|
signal_handlers[signal.SIGINT]()
|
||||||
|
|
||||||
|
_FakeCrabstero.login_callback = request_shutdown
|
||||||
|
|
||||||
|
assert cli.main() == 130
|
||||||
|
assert ("notify", "READY=1") not in events
|
||||||
|
assert ("notify", "STOPPING=1") in events
|
||||||
|
assert "close" in events
|
||||||
|
assert signal_handlers == {}
|
||||||
|
|
||||||
|
def test_keyboard_interrupt_returns_130(
|
||||||
|
self,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
"""KeyboardInterrupt maps to the conventional interrupted exit code."""
|
||||||
|
|
||||||
|
def interrupt(coro: Coroutine[Any, Any, None]) -> None:
|
||||||
|
coro.close()
|
||||||
|
raise KeyboardInterrupt
|
||||||
|
|
||||||
|
monkeypatch.setattr(sys, "argv", ["crabstero", "--token", "cli-token"])
|
||||||
|
monkeypatch.setattr("crabstero.cli.uvloop.run", interrupt)
|
||||||
|
|
||||||
|
assert cli.main() == 130
|
||||||
@@ -0,0 +1,164 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Unit tests for the flag convenience wrappers and enums.
|
||||||
|
|
||||||
|
Tests cover _entity_id, Flag and EntityType enums, and the high-level
|
||||||
|
set_flag/clear_flag/is_flag_set wrappers from crabstero.flags without crossing
|
||||||
|
the SQLite database boundary.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero.database import Database
|
||||||
|
from crabstero.flags import (
|
||||||
|
EntityType,
|
||||||
|
Flag,
|
||||||
|
_entity_id,
|
||||||
|
clear_flag,
|
||||||
|
is_flag_set,
|
||||||
|
set_flag,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _StubSnowflake:
|
||||||
|
"""Minimal stand-in for objects that expose a Discord snowflake ID."""
|
||||||
|
|
||||||
|
def __init__(self, *, entity_id: int = 123456789) -> None:
|
||||||
|
self.id = entity_id
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeFlagDatabase(Database):
|
||||||
|
"""In-memory flag store that records wrapper calls below SQLite."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.flags: set[tuple[str, str, str]] = set()
|
||||||
|
|
||||||
|
async def set_flag(self, entity_type: str, entity_id: str, flag_name: str) -> None:
|
||||||
|
"""Record a flag set operation."""
|
||||||
|
self.flags.add((entity_type, entity_id, flag_name))
|
||||||
|
|
||||||
|
async def clear_flag(
|
||||||
|
self,
|
||||||
|
entity_type: str,
|
||||||
|
entity_id: str,
|
||||||
|
flag_name: str,
|
||||||
|
) -> None:
|
||||||
|
"""Record a flag clear operation."""
|
||||||
|
self.flags.discard((entity_type, entity_id, flag_name))
|
||||||
|
|
||||||
|
async def is_flag_set(
|
||||||
|
self,
|
||||||
|
entity_type: str,
|
||||||
|
entity_id: str,
|
||||||
|
flag_name: str,
|
||||||
|
) -> bool:
|
||||||
|
"""Return whether a flag has been set."""
|
||||||
|
return (entity_type, entity_id, flag_name) in self.flags
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def flag_db() -> _FakeFlagDatabase:
|
||||||
|
"""Return an isolated flag store for wrapper tests."""
|
||||||
|
return _FakeFlagDatabase()
|
||||||
|
|
||||||
|
|
||||||
|
class TestEntityId:
|
||||||
|
"""Entity ID extraction from Discord objects and integers."""
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("entity", "expected"),
|
||||||
|
[
|
||||||
|
pytest.param(42, "42", id="raw-integer"),
|
||||||
|
pytest.param(
|
||||||
|
_StubSnowflake(entity_id=99),
|
||||||
|
"99",
|
||||||
|
id="snowflake-object",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_extraction(self, entity: object, expected: str) -> None:
|
||||||
|
"""Converts the entity to the expected string ID."""
|
||||||
|
assert _entity_id(entity) == expected # type: ignore[arg-type]
|
||||||
|
|
||||||
|
|
||||||
|
class TestFlagEnums:
|
||||||
|
"""Flag and EntityType enum values."""
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("member", "expected"),
|
||||||
|
[
|
||||||
|
pytest.param(Flag.NO_REPLY, "noReply", id="no-reply"),
|
||||||
|
pytest.param(Flag.NO_INGEST, "noIngest", id="no-ingest"),
|
||||||
|
pytest.param(Flag.ALLOW_PINGS, "allowPings", id="allow-pings"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_flag_values(self, member: Flag, expected: str) -> None:
|
||||||
|
"""Flag enum value matches the expected database string."""
|
||||||
|
assert member.value == expected
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("member", "expected"),
|
||||||
|
[
|
||||||
|
pytest.param(EntityType.CHANNEL, "channel", id="channel"),
|
||||||
|
pytest.param(EntityType.SERVER, "server", id="server"),
|
||||||
|
pytest.param(EntityType.USER, "user", id="user"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_entity_type_values(self, member: EntityType, expected: str) -> None:
|
||||||
|
"""EntityType enum value matches the expected database string."""
|
||||||
|
assert member.value == expected
|
||||||
|
|
||||||
|
|
||||||
|
class TestSetClearCheck:
|
||||||
|
"""High-level flag set/clear/check cycle through the flags module."""
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"flag",
|
||||||
|
[
|
||||||
|
pytest.param(Flag.NO_REPLY, id="no-reply"),
|
||||||
|
pytest.param(Flag.NO_INGEST, id="no-ingest"),
|
||||||
|
pytest.param(Flag.ALLOW_PINGS, id="allow-pings"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_set_then_check(self, flag_db: _FakeFlagDatabase, flag: Flag) -> None:
|
||||||
|
"""A set flag is reported as set."""
|
||||||
|
await set_flag(flag_db, 1, EntityType.CHANNEL, flag)
|
||||||
|
assert await is_flag_set(flag_db, 1, EntityType.CHANNEL, flag) is True
|
||||||
|
|
||||||
|
async def test_unset_returns_false(self, flag_db: _FakeFlagDatabase) -> None:
|
||||||
|
"""An unset flag is reported as not set."""
|
||||||
|
assert (
|
||||||
|
await is_flag_set(
|
||||||
|
flag_db,
|
||||||
|
1,
|
||||||
|
EntityType.CHANNEL,
|
||||||
|
Flag.NO_REPLY,
|
||||||
|
)
|
||||||
|
is False
|
||||||
|
)
|
||||||
|
|
||||||
|
async def test_clear_removes_flag(self, flag_db: _FakeFlagDatabase) -> None:
|
||||||
|
"""A cleared flag is no longer reported as set."""
|
||||||
|
await set_flag(flag_db, 1, EntityType.CHANNEL, Flag.NO_REPLY)
|
||||||
|
await clear_flag(flag_db, 1, EntityType.CHANNEL, Flag.NO_REPLY)
|
||||||
|
assert (
|
||||||
|
await is_flag_set(
|
||||||
|
flag_db,
|
||||||
|
1,
|
||||||
|
EntityType.CHANNEL,
|
||||||
|
Flag.NO_REPLY,
|
||||||
|
)
|
||||||
|
is False
|
||||||
|
)
|
||||||
@@ -0,0 +1,470 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Unit tests for Markov chain ingestion and generation.
|
||||||
|
|
||||||
|
Tests cover is_complete_sentence, _ingest_sentence, ingest, and generate
|
||||||
|
from crabstero.markov without crossing the SQLite database boundary.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from crabstero.database import Database, StartWord, Transition
|
||||||
|
from crabstero.markov import (
|
||||||
|
DEFAULT_SENTENCE_END,
|
||||||
|
_ingest_sentence,
|
||||||
|
_uningest_sentence,
|
||||||
|
generate,
|
||||||
|
ingest,
|
||||||
|
is_complete_sentence,
|
||||||
|
uningest,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeMarkovDatabase(Database):
|
||||||
|
"""In-memory Markov store with deterministic first-row selection.
|
||||||
|
|
||||||
|
The fake preserves one-occurrence-at-a-time add/remove semantics so
|
||||||
|
ingestion and uningestion tests can stay below the SQLite boundary.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.start_words: list[StartWord] = []
|
||||||
|
self.transitions: list[Transition] = []
|
||||||
|
|
||||||
|
async def add_markov_data(
|
||||||
|
self,
|
||||||
|
start_words: list[StartWord],
|
||||||
|
transitions: list[Transition],
|
||||||
|
) -> None:
|
||||||
|
"""Record Markov rows exactly as the application would write them."""
|
||||||
|
self.start_words.extend(start_words)
|
||||||
|
self.transitions.extend(transitions)
|
||||||
|
|
||||||
|
async def remove_markov_data(
|
||||||
|
self,
|
||||||
|
start_words: list[StartWord],
|
||||||
|
transitions: list[Transition],
|
||||||
|
) -> None:
|
||||||
|
"""Remove at most one matching row per requested Markov row."""
|
||||||
|
for start_word in start_words:
|
||||||
|
with contextlib.suppress(ValueError):
|
||||||
|
self.start_words.remove(start_word)
|
||||||
|
for transition in transitions:
|
||||||
|
with contextlib.suppress(ValueError):
|
||||||
|
self.transitions.remove(transition)
|
||||||
|
|
||||||
|
async def get_random_start_word(self, channel_id: int) -> str | None:
|
||||||
|
"""Return the first start word for deterministic generation tests."""
|
||||||
|
return next(
|
||||||
|
(row.word for row in self.start_words if row.channel_id == channel_id),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def get_random_next_word(self, channel_id: int, word: str) -> str | None:
|
||||||
|
"""Return the first matching transition for deterministic tests."""
|
||||||
|
return next(
|
||||||
|
(
|
||||||
|
row.next_word
|
||||||
|
for row in self.transitions
|
||||||
|
if row.channel_id == channel_id and row.word == word
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def get_random_completing_next_word(
|
||||||
|
self,
|
||||||
|
channel_id: int,
|
||||||
|
word: str,
|
||||||
|
) -> str | None:
|
||||||
|
"""Return the first sentence-ending transition for deterministic tests."""
|
||||||
|
return next(
|
||||||
|
(
|
||||||
|
row.next_word
|
||||||
|
for row in self.transitions
|
||||||
|
if row.channel_id == channel_id
|
||||||
|
and row.word == word
|
||||||
|
and is_complete_sentence(row.next_word)
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def markov_db() -> _FakeMarkovDatabase:
|
||||||
|
"""Return an isolated deterministic Markov store for one test."""
|
||||||
|
return _FakeMarkovDatabase()
|
||||||
|
|
||||||
|
|
||||||
|
class TestIsCompleteSentence:
|
||||||
|
"""Sentence completeness detection."""
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("sentence", "expected"),
|
||||||
|
[
|
||||||
|
pytest.param("Hello.", True, id="period"),
|
||||||
|
pytest.param("Wow!", True, id="exclamation"),
|
||||||
|
pytest.param("Really?", True, id="question"),
|
||||||
|
pytest.param(f"Hello{DEFAULT_SENTENCE_END}", True, id="section-sign"),
|
||||||
|
pytest.param("Hello", False, id="no-punctuation"),
|
||||||
|
pytest.param("", False, id="empty-string"),
|
||||||
|
pytest.param("Hello. ", False, id="trailing-space"),
|
||||||
|
pytest.param("Hello,", False, id="comma"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_detection(self, sentence: str, expected: bool) -> None: # noqa: FBT001
|
||||||
|
"""Correctly identifies sentence completeness."""
|
||||||
|
assert is_complete_sentence(sentence) is expected
|
||||||
|
|
||||||
|
|
||||||
|
class TestIngestSentence:
|
||||||
|
"""Single sentence ingestion into the Markov chain."""
|
||||||
|
|
||||||
|
async def test_stores_start_word(self, markov_db: _FakeMarkovDatabase) -> None:
|
||||||
|
"""First word of the sentence is stored as a start word."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Hello world.")
|
||||||
|
|
||||||
|
assert markov_db.start_words == [StartWord(1, 100, "Hello")]
|
||||||
|
|
||||||
|
async def test_stores_transitions(self, markov_db: _FakeMarkovDatabase) -> None:
|
||||||
|
"""Adjacent words create transitions."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Hello world.")
|
||||||
|
|
||||||
|
assert markov_db.transitions == [Transition(1, 100, "Hello", "world.")]
|
||||||
|
|
||||||
|
async def test_appends_sentence_end_if_missing(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Unpunctuated sentence gets the default sentence-end marker."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Hello world")
|
||||||
|
|
||||||
|
assert markov_db.transitions == [
|
||||||
|
Transition(1, 100, "Hello", f"world{DEFAULT_SENTENCE_END}"),
|
||||||
|
]
|
||||||
|
|
||||||
|
async def test_preserves_existing_punctuation(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Already-punctuated sentence keeps its terminator."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Hello world!")
|
||||||
|
|
||||||
|
assert markov_db.transitions == [Transition(1, 100, "Hello", "world!")]
|
||||||
|
|
||||||
|
async def test_single_word_stores_nothing(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""A single-word sentence produces no start words or transitions."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Hello.")
|
||||||
|
|
||||||
|
assert markov_db.start_words == []
|
||||||
|
assert markov_db.transitions == []
|
||||||
|
|
||||||
|
async def test_stores_all_transitions(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""All adjacent word pairs create transitions."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "A B C.")
|
||||||
|
|
||||||
|
assert markov_db.transitions == [
|
||||||
|
Transition(1, 100, "A", "B"),
|
||||||
|
Transition(1, 100, "B", "C."),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestIngestParagraph:
|
||||||
|
"""Paragraph ingestion splits into sentences."""
|
||||||
|
|
||||||
|
async def test_single_sentence(self, markov_db: _FakeMarkovDatabase) -> None:
|
||||||
|
"""A single sentence paragraph is ingested."""
|
||||||
|
await ingest(markov_db, 1, 100, "Hello world.")
|
||||||
|
|
||||||
|
assert markov_db.start_words == [StartWord(1, 100, "Hello")]
|
||||||
|
|
||||||
|
async def test_multiple_sentences(self, markov_db: _FakeMarkovDatabase) -> None:
|
||||||
|
"""Multiple sentences are split and ingested individually."""
|
||||||
|
await ingest(markov_db, 1, 100, "Hello world. Goodbye world!")
|
||||||
|
|
||||||
|
assert {row.word for row in markov_db.start_words} == {"Hello", "Goodbye"}
|
||||||
|
|
||||||
|
async def test_normalizes_whitespace(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Extra spaces and newlines are collapsed."""
|
||||||
|
await ingest(markov_db, 1, 100, "Hello world.\nGoodbye world!")
|
||||||
|
|
||||||
|
assert Transition(1, 100, "Hello", "world.") in markov_db.transitions
|
||||||
|
assert Transition(1, 100, "Goodbye", "world!") in markov_db.transitions
|
||||||
|
|
||||||
|
async def test_appends_default_end(self, markov_db: _FakeMarkovDatabase) -> None:
|
||||||
|
"""Unpunctuated paragraph gets the default sentence-end marker."""
|
||||||
|
await ingest(markov_db, 1, 100, "Hello world")
|
||||||
|
|
||||||
|
assert markov_db.transitions == [
|
||||||
|
Transition(1, 100, "Hello", f"world{DEFAULT_SENTENCE_END}"),
|
||||||
|
]
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"paragraph",
|
||||||
|
[
|
||||||
|
pytest.param("", id="empty-string"),
|
||||||
|
pytest.param(" ", id="whitespace-only"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_empty_input_stores_nothing(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
paragraph: str,
|
||||||
|
) -> None:
|
||||||
|
"""Empty or whitespace-only input does not store any data."""
|
||||||
|
await ingest(markov_db, 1, 100, paragraph)
|
||||||
|
|
||||||
|
assert markov_db.start_words == []
|
||||||
|
assert markov_db.transitions == []
|
||||||
|
|
||||||
|
async def test_splits_on_punctuation_followed_by_space(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Punctuation followed by a space splits into separate sentences."""
|
||||||
|
await ingest(markov_db, 1, 100, "Dr. Smith likes cats")
|
||||||
|
|
||||||
|
# "Dr." splits off as a single-word sentence (stores nothing).
|
||||||
|
# "Smith likes cats" becomes a sentence with "Smith" as start word.
|
||||||
|
assert markov_db.start_words == [StartWord(1, 100, "Smith")]
|
||||||
|
|
||||||
|
async def test_no_split_without_space_after_punctuation(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Punctuation not followed by a space keeps words together."""
|
||||||
|
await ingest(markov_db, 1, 100, "Hello.World is here")
|
||||||
|
|
||||||
|
assert markov_db.start_words == [StartWord(1, 100, "Hello.World")]
|
||||||
|
|
||||||
|
|
||||||
|
class TestGenerate:
|
||||||
|
"""Markov chain text generation."""
|
||||||
|
|
||||||
|
async def test_fallback_on_empty_channel(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Returns an informational message when the channel has no data."""
|
||||||
|
result = await generate(markov_db, 1)
|
||||||
|
expected = (
|
||||||
|
"I do not have enough data to generate a message yet."
|
||||||
|
" Chat a bit more so I can learn how this channel talks."
|
||||||
|
)
|
||||||
|
assert result == expected
|
||||||
|
|
||||||
|
async def test_generates_from_ingested_data(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Generated text uses words from ingested data."""
|
||||||
|
await ingest(markov_db, 1, 100, "The quick brown fox.")
|
||||||
|
|
||||||
|
result = await generate(markov_db, 1)
|
||||||
|
assert result == "The quick brown fox."
|
||||||
|
|
||||||
|
async def test_strips_section_sign(self, markov_db: _FakeMarkovDatabase) -> None:
|
||||||
|
"""The internal section sign marker never appears in output."""
|
||||||
|
await ingest(markov_db, 1, 100, "Hello world")
|
||||||
|
|
||||||
|
result = await generate(markov_db, 1)
|
||||||
|
assert result == "Hello world"
|
||||||
|
|
||||||
|
async def test_respects_hard_limit(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Output is truncated at hard_limit."""
|
||||||
|
words = [f"w{i}" for i in range(100)]
|
||||||
|
text = " ".join(words)
|
||||||
|
await ingest(markov_db, 1, 100, text)
|
||||||
|
|
||||||
|
result = await generate(markov_db, 1, soft_limit=10, hard_limit=49)
|
||||||
|
assert result == "w0 w1 w2 w3 w4 w5 w6 w7 w8 w9 w10 w11 w12 w13 w14"
|
||||||
|
|
||||||
|
async def test_prefers_completing_word_after_soft_limit(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""After soft_limit, generation prefers completing words."""
|
||||||
|
# Chain: A -> B -> C -> D -> {E, "end."}
|
||||||
|
# A, B, C have only non-completing transitions, so past the soft limit
|
||||||
|
# the loop falls back to get_random_next_word for each. D has both a
|
||||||
|
# continuing ("E") and completing ("end.") transition, so
|
||||||
|
# get_random_completing_next_word deterministically picks "end.".
|
||||||
|
await markov_db.add_markov_data(
|
||||||
|
[StartWord(1, 100, "A")],
|
||||||
|
[
|
||||||
|
Transition(1, 100, "A", "B"),
|
||||||
|
Transition(1, 100, "B", "C"),
|
||||||
|
Transition(1, 100, "C", "D"),
|
||||||
|
Transition(1, 100, "D", "E"),
|
||||||
|
Transition(1, 100, "D", "F"),
|
||||||
|
Transition(1, 100, "D", "G"),
|
||||||
|
Transition(1, 100, "D", "H"),
|
||||||
|
Transition(1, 100, "D", "end."),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await generate(markov_db, 1, soft_limit=1, hard_limit=1000)
|
||||||
|
assert result == "A B C D end."
|
||||||
|
|
||||||
|
async def test_start_word_already_ends_sentence(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Generation stops immediately when the start word is sentence-ending."""
|
||||||
|
await markov_db.add_markov_data([StartWord(1, 100, "Yes.")], [])
|
||||||
|
|
||||||
|
result = await generate(markov_db, 1)
|
||||||
|
assert result == "Yes."
|
||||||
|
|
||||||
|
async def test_chain_dead_end(self, markov_db: _FakeMarkovDatabase) -> None:
|
||||||
|
"""Generation stops when no next word exists (dead-end chain)."""
|
||||||
|
await markov_db.add_markov_data([StartWord(1, 100, "Hello")], [])
|
||||||
|
|
||||||
|
result = await generate(markov_db, 1)
|
||||||
|
assert result == "Hello"
|
||||||
|
|
||||||
|
async def test_hard_limit_strips_section_sign(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Section sign at the truncation boundary is stripped."""
|
||||||
|
await markov_db.add_markov_data(
|
||||||
|
[StartWord(1, 100, "A")],
|
||||||
|
[
|
||||||
|
Transition(1, 100, "A", "B"),
|
||||||
|
Transition(1, 100, "B", f"C{DEFAULT_SENTENCE_END}"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await generate(markov_db, 1, soft_limit=100, hard_limit=6)
|
||||||
|
assert result == "A B C"
|
||||||
|
|
||||||
|
async def test_soft_limit_falls_back_to_regular_next_word(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""After soft_limit, falls back when no completing word exists."""
|
||||||
|
await markov_db.add_markov_data(
|
||||||
|
[StartWord(1, 100, "A")],
|
||||||
|
[Transition(1, 100, "A", "B")],
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await generate(markov_db, 1, soft_limit=1, hard_limit=1000)
|
||||||
|
assert result == "A B"
|
||||||
|
|
||||||
|
async def test_hard_limit_truncates_mid_word(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Hard limit slices output even when it falls inside a word."""
|
||||||
|
await markov_db.add_markov_data(
|
||||||
|
[StartWord(1, 100, "AB")],
|
||||||
|
[Transition(1, 100, "AB", "CDEF")],
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await generate(markov_db, 1, soft_limit=100, hard_limit=5)
|
||||||
|
assert result == "AB CD"
|
||||||
|
|
||||||
|
async def test_channel_isolation(self, markov_db: _FakeMarkovDatabase) -> None:
|
||||||
|
"""Data ingested into one channel does not leak into another."""
|
||||||
|
await ingest(markov_db, 1, 100, "Channel one data.")
|
||||||
|
await ingest(markov_db, 2, 100, "Channel two data.")
|
||||||
|
|
||||||
|
result = await generate(markov_db, 3)
|
||||||
|
expected = (
|
||||||
|
"I do not have enough data to generate a message yet."
|
||||||
|
" Chat a bit more so I can learn how this channel talks."
|
||||||
|
)
|
||||||
|
assert result == expected
|
||||||
|
|
||||||
|
|
||||||
|
class TestUningestSentence:
|
||||||
|
"""Single sentence uningest from the Markov chain."""
|
||||||
|
|
||||||
|
async def test_removes_start_word_and_transition(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Uningest removes the start word and transition added by ingest."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Hello world.")
|
||||||
|
await _uningest_sentence(markov_db, 1, 100, "Hello world.")
|
||||||
|
assert markov_db.start_words == []
|
||||||
|
assert markov_db.transitions == []
|
||||||
|
|
||||||
|
async def test_preserves_other_data(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Uningest only removes data for the specified sentence."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Hello world.")
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Goodbye world.")
|
||||||
|
await _uningest_sentence(markov_db, 1, 100, "Hello world.")
|
||||||
|
assert markov_db.start_words == [StartWord(1, 100, "Goodbye")]
|
||||||
|
assert markov_db.transitions == [Transition(1, 100, "Goodbye", "world.")]
|
||||||
|
|
||||||
|
async def test_handles_missing_punctuation(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Uningest appends the default sentence end, matching ingest behavior."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Hello world")
|
||||||
|
await _uningest_sentence(markov_db, 1, 100, "Hello world")
|
||||||
|
assert markov_db.start_words == []
|
||||||
|
assert markov_db.transitions == []
|
||||||
|
|
||||||
|
async def test_single_word_is_noop(self, markov_db: _FakeMarkovDatabase) -> None:
|
||||||
|
"""Uningesting a single-word sentence does not error."""
|
||||||
|
await _ingest_sentence(markov_db, 1, 100, "Hello.")
|
||||||
|
await _uningest_sentence(markov_db, 1, 100, "Hello.")
|
||||||
|
assert markov_db.start_words == []
|
||||||
|
assert markov_db.transitions == []
|
||||||
|
|
||||||
|
|
||||||
|
class TestUningestParagraph:
|
||||||
|
"""Paragraph-level uningest."""
|
||||||
|
|
||||||
|
async def test_uningest_multiple_sentences(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Uningest reverses a multi-sentence paragraph."""
|
||||||
|
await ingest(markov_db, 1, 100, "Hello world. Goodbye world!")
|
||||||
|
await uningest(markov_db, 1, 100, "Hello world. Goodbye world!")
|
||||||
|
assert markov_db.start_words == []
|
||||||
|
assert markov_db.transitions == []
|
||||||
|
|
||||||
|
async def test_uningest_preserves_duplicate_data(
|
||||||
|
self,
|
||||||
|
markov_db: _FakeMarkovDatabase,
|
||||||
|
) -> None:
|
||||||
|
"""Uningesting one copy leaves the other intact."""
|
||||||
|
await ingest(markov_db, 1, 100, "Hello world.")
|
||||||
|
await ingest(markov_db, 1, 100, "Hello world.")
|
||||||
|
await uningest(markov_db, 1, 100, "Hello world.")
|
||||||
|
assert markov_db.start_words == [StartWord(1, 100, "Hello")]
|
||||||
|
assert markov_db.transitions == [Transition(1, 100, "Hello", "world.")]
|
||||||
@@ -0,0 +1,74 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Unit tests for the Prometheus metrics module.
|
||||||
|
|
||||||
|
Tests cover metric object registration.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from prometheus_client import generate_latest
|
||||||
|
|
||||||
|
# Importing the module registers its Prometheus metrics with the default registry.
|
||||||
|
import crabstero.metrics # noqa: F401 - imported for registration side effects
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="module")
|
||||||
|
def prometheus_output() -> str:
|
||||||
|
"""Return Prometheus exposition text after module-level metric registration."""
|
||||||
|
return generate_latest().decode()
|
||||||
|
|
||||||
|
|
||||||
|
class TestMetricObjects:
|
||||||
|
"""All expected metric objects exist and are registered."""
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"name",
|
||||||
|
[
|
||||||
|
pytest.param("crabstero_build_info", id="build-info"),
|
||||||
|
pytest.param("crabstero_replies_sent_total", id="replies-sent"),
|
||||||
|
pytest.param("crabstero_embeds_generated_total", id="embeds-generated"),
|
||||||
|
pytest.param("crabstero_messages_processed_total", id="messages-processed"),
|
||||||
|
pytest.param("crabstero_pingme_total", id="pingme"),
|
||||||
|
pytest.param("crabstero_forgetme_total", id="forgetme"),
|
||||||
|
pytest.param("crabstero_errors_total", id="errors"),
|
||||||
|
pytest.param("crabstero_discord_latency_seconds", id="discord-latency"),
|
||||||
|
pytest.param("crabstero_guild_count", id="guild-count"),
|
||||||
|
pytest.param("crabstero_ingestion_active_channels", id="ingestion-active"),
|
||||||
|
pytest.param(
|
||||||
|
"crabstero_message_ingestion_duration_seconds",
|
||||||
|
id="message-ingestion-duration",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
"crabstero_channel_ingestion_duration_seconds",
|
||||||
|
id="channel-ingestion-duration",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
"crabstero_generation_duration_seconds",
|
||||||
|
id="generation-duration",
|
||||||
|
),
|
||||||
|
pytest.param(
|
||||||
|
"crabstero_channel_ingestion_messages",
|
||||||
|
id="channel-ingestion-messages",
|
||||||
|
),
|
||||||
|
pytest.param("crabstero_discord_events_total", id="discord-events"),
|
||||||
|
pytest.param(
|
||||||
|
"crabstero_messages_uningested_total",
|
||||||
|
id="messages-uningested",
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_metric_in_output(self, prometheus_output: str, name: str) -> None:
|
||||||
|
"""Each declared metric appears in the Prometheus text output."""
|
||||||
|
assert name in prometheus_output, f"{name} not found in Prometheus output"
|
||||||
@@ -0,0 +1,26 @@
|
|||||||
|
# Copyright 2026 Logan Fick
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
"""Unit tests for package version availability."""
|
||||||
|
|
||||||
|
import crabstero
|
||||||
|
|
||||||
|
|
||||||
|
class TestVersion:
|
||||||
|
"""Package version metadata."""
|
||||||
|
|
||||||
|
def test_version_is_available(self) -> None:
|
||||||
|
"""__version__ is a non-empty string."""
|
||||||
|
assert isinstance(crabstero.__version__, str)
|
||||||
|
assert crabstero.__version__
|
||||||
Reference in New Issue
Block a user