59 Commits
Author SHA1 Message Date
LogalDeveloper d83740e40f Updated dependencies.
Audit / Dependencies (push) Successful in 7s
CD / Publish (push) Successful in 5s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 22s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-07-05 15:36:32 -04:00
LogalDeveloper 8e742426c9 Added README.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-07-05 15:26:51 -04:00
LogalDeveloper 29edbdba18 Improved shutdown cleanup ordering.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 6s
2026-06-30 09:08:35 -04:00
LogalDeveloper 73c103c486 Capped generated embed titles at Discord's limit.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 4s
2026-06-29 20:25:36 -04:00
LogalDeveloper 7cdb27f69c Fixed Ruff and Mypy issues.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 8s
CI / Tests (push) Successful in 36s
2026-06-29 19:35:08 -04:00
LogalDeveloper 0821169a3f Updated dependencies.
Audit / Dependencies (push) Successful in 8s
CI / Formatting (push) Failing after 4s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Failing after 10s
CI / Spelling (push) Successful in 5s
2026-06-29 15:36:16 -04:00
LogalDeveloper 1dd9db9151 Reworked Markov storage to reduce database size.
CI / Formatting (push) Failing after 38s
CI / Linting (push) Successful in 8s
CI / Tests (push) Successful in 35s
CI / Type Checking (push) Failing after 12s
CI / Spelling (push) Successful in 7s
2026-06-29 12:29:19 -04:00
LogalDeveloper d23c11cf11 Updated dependencies.
Audit / Dependencies (push) Successful in 7s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 39s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-06-18 19:56:26 -04:00
LogalDeveloper 5bbb0bbde5 Expanded integration coverage and enforced test categories.
Audit / Dependencies (push) Successful in 9s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 1m7s
CI / Type Checking (push) Successful in 12s
CI / Spelling (push) Successful in 8s
2026-06-18 16:10:40 -04:00
LogalDeveloper 49062159f9 Moved project into a src-based layout and reorganized tests.
CI / Formatting (push) Failing after 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 8s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 4s
2026-06-16 14:21:02 -04:00
LogalDeveloper 1aafecf20e Separated CLI orchestration from the Crabstero bot object.
CI / Formatting (push) Failing after 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 8s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-06-15 20:35:02 -04:00
LogalDeveloper 68f000ef27 Added configurable permissions for metrics Unix sockets.
CI / Formatting (push) Successful in 29s
CI / Linting (push) Successful in 7s
CI / Tests (push) Successful in 11s
CI / Type Checking (push) Successful in 12s
CI / Spelling (push) Successful in 8s
2026-06-11 19:40:53 -04:00
LogalDeveloper 6d7f26f7be Addressed Ruff checks for Unix socket startup path handling.
Audit / Dependencies (push) Successful in 6s
CD / Publish (push) Successful in 5s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-06-11 15:11:51 -04:00
LogalDeveloper 446c0cd9ca Added Unix socket support for the Prometheus metrics server.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Failing after 5s
CI / Tests (push) Successful in 8s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-06-11 15:07:44 -04:00
LogalDeveloper 7b395e1834 Added systemd notify support with ready, stopping, and watchdog notifications.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-06-11 13:55:04 -04:00
LogalDeveloper b1625bfab5 Updated dependencies.
Audit / Dependencies (push) Successful in 8s
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 6s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-06-11 13:34:45 -04:00
LogalDeveloper b22b1ffab5 Aligned pyproject.toml with standard tooling conventions.
CI / Formatting (push) Successful in 29s
CI / Linting (push) Successful in 6s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-06-11 13:28:18 -04:00
LogalDeveloper de19074fc0 Switched Syft SBOM generation to runtime-only venv scan.
Audit / Dependencies (push) Successful in 7s
CD / Publish (push) Successful in 5s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 4s
2026-05-15 20:33:39 -04:00
LogalDeveloper ef0d1ce144 Updated dependencies.
Audit / Dependencies (push) Successful in 8s
CD / Publish (push) Successful in 6s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 9s
CI / Tests (push) Successful in 12s
CI / Type Checking (push) Successful in 12s
CI / Spelling (push) Successful in 8s
2026-05-15 12:44:52 -04:00
LogalDeveloper 85c71faf2d Added Syft SBOM generation to CD workflow.
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 8s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-05-15 12:35:51 -04:00
LogalDeveloper 362541e13b Updated Gitea Actions workflows to use new CI image.
CI / Formatting (push) Successful in 36s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 5s
2026-05-15 12:13:43 -04:00
LogalDeveloper 920f9f9f2d Updated dependencies.
Audit / Dependencies (push) Successful in 7s
CD / Publish (push) Successful in 5s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 11s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 11s
2026-04-29 09:47:41 -04:00
LogalDeveloper 23172111e9 Hardened Gitea Actions workflows and updated action pins.
CI / Formatting (push) Successful in 26s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 5s
2026-04-29 09:37:17 -04:00
LogalDeveloper b680f52126 Scoped GITEA_TOKEN to least privilege in Gitea Actions workflows.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 5s
2026-04-22 12:22:10 -04:00
LogalDeveloper eb81fa478d Updated dependencies.
Audit / Dependencies (push) Successful in 23s
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 7s
2026-04-21 08:49:52 -04:00
LogalDeveloper 68975fcadd Switched test database fixture to in-memory SQLite.
CD / Publish (push) Successful in 6s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 7s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 4s
Audit / Dependencies (push) Failing after 6s
2026-04-05 09:49:35 -04:00
LogalDeveloper 0b6e292cbd Updated dependencies.
Audit / Dependencies (push) Successful in 7s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 29s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-04-05 09:46:59 -04:00
LogalDeveloper 68e0bc190c Switched event loop to uvloop for improved async performance.
Audit / Dependencies (push) Successful in 7s
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 28s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 4s
2026-04-05 09:44:02 -04:00
LogalDeveloper c577dc5267 Updated dependencies.
Audit / Dependencies (push) Successful in 6s
CD / Publish (push) Successful in 5s
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 27s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-29 20:17:09 -04:00
LogalDeveloper abe1444e4c Enabled all Ruff linter rules and fixed resulting violations.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 28s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-29 20:08:30 -04:00
LogalDeveloper 603a927f12 Replaced raw string literals with Flag and EntityType enums in database tests.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 29s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
Audit / Dependencies (push) Failing after 7s
2026-03-24 19:18:41 -04:00
LogalDeveloper 2dc8206f6f Added scoping and isolation tests for database operations and uningest.
Audit / Dependencies (push) Successful in 16s
CD / Publish (push) Successful in 7s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 6s
CI / Tests (push) Successful in 34s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-03-23 20:26:25 -04:00
LogalDeveloper 2b65dfe076 Pre-initialized labels for slash command metrics for immediate visibility.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-03-23 17:44:54 -04:00
LogalDeveloper 0013726b7f Fixed Ruff formatting issue.
Audit / Dependencies (push) Successful in 6s
CD / Publish (push) Successful in 5s
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-23 16:20:45 -04:00
LogalDeveloper b8f509ac63 Updated dependencies.
Audit / Dependencies (push) Successful in 7s
CI / Formatting (push) Failing after 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 24s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 4s
2026-03-23 16:17:07 -04:00
LogalDeveloper b0fdf12808 Standardized user-facing messages.
CI / Formatting (push) Failing after 4s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 14s
CI / Spelling (push) Successful in 9s
2026-03-23 16:02:54 -04:00
LogalDeveloper cf46b71883 Added integration tests for uningest_message.
CI / Formatting (push) Successful in 6s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 24s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-23 15:45:31 -04:00
LogalDeveloper bc9f7fa539 Removed redundant metrics, fixed description accuracy, and excluded no-data fallback from generation timing.
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-03-23 15:31:43 -04:00
LogalDeveloper 70095f3184 Replaced regex with native str methods for whitespace normalization and mention extraction.
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-03-23 14:39:38 -04:00
LogalDeveloper 0fa51d6199 Added /forgetme slash command with atomic user data deletion, confirmation UI, and Prometheus metrics.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-23 10:06:55 -04:00
LogalDeveloper ca3011011a Added async lock to database transactions to prevent concurrent transaction conflicts.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 22s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-23 09:49:44 -04:00
LogalDeveloper 53b8bd3fae Switched to explicit SQLite transactions with rollback safety and batched delete operations.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 23s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-23 09:16:46 -04:00
LogalDeveloper 093426ac61 Refactored ingestion to semaphore-bounded tasks, centralized error handling, and consolidated metrics.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 22s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-22 22:45:31 -04:00
LogalDeveloper c1f967ece9 Fixed cache shutdown to properly await task cancellation and added names to background tasks.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 22s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-03-22 20:17:06 -04:00
LogalDeveloper 23cc721612 Modernized codebase with NamedTuples, StrEnum, override decorators, slots, and other idiomatic improvements.
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 21s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-03-22 20:02:31 -04:00
LogalDeveloper a23d252f27 Fixed mypy override compatibility error in bot start method signature.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 21s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 4s
Audit / Dependencies (push) Successful in 6s
2026-03-22 14:54:12 -04:00
LogalDeveloper 2cf00fde4f Improved code quality with more idiomatic Python patterns and safer initialization.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 21s
CI / Type Checking (push) Failing after 10s
CI / Spelling (push) Successful in 5s
2026-03-22 14:41:38 -04:00
LogalDeveloper 371f7d2ca3 Added message uningest to reverse Markov data when recently ingested messages are deleted.
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 20s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-22 14:05:58 -04:00
LogalDeveloper 12e5e92fe4 Updated dependencies.
Audit / Dependencies (push) Successful in 7s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 16s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-19 09:01:33 -04:00
LogalDeveloper a7b4337835 Replaced fallback database seeding with an informational message for empty channels.
Audit / Dependencies (push) Successful in 6s
CD / Publish (push) Successful in 5s
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 16s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-18 20:01:12 -04:00
LogalDeveloper a8d3a5c421 Added Prometheus metrics support with opt-in CLI flag.
Audit / Dependencies (push) Successful in 7s
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 16s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-18 19:43:30 -04:00
LogalDeveloper 6aa8527adc Simplified database layer and consolidated Markov writes into a single transaction.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 16s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s
2026-03-17 17:07:28 -04:00
LogalDeveloper 56d3c36c7d Replaced redundant database query with known fallback start word to fix mypy type narrowing.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 9s
CI / Type Checking (push) Successful in 11s
CI / Spelling (push) Successful in 5s
2026-03-17 14:28:47 -04:00
LogalDeveloper 2d13dfb063 Added unit tests for database, markov, flags, CLI, and version modules.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 9s
CI / Type Checking (push) Failing after 10s
CI / Spelling (push) Successful in 5s
2026-03-17 14:13:31 -04:00
LogalDeveloper 34ae3b3e64 Aligned CI workflows and pyproject.toml with cross-project infrastructure conventions.
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 4s
CI / Tests (push) Successful in 5s
CI / Type Checking (push) Successful in 9s
CI / Spelling (push) Successful in 4s
Audit / Dependencies (push) Successful in 8s
2026-03-14 17:45:07 -04:00
LogalDeveloper 51f33a8dac Deduplicated allowed mention user IDs to prevent Discord API rejection.
CI / Formatting (push) Successful in 4s
CI / Linting (push) Successful in 5s
CI / Tests (push) Successful in 6s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 4s
2026-03-14 17:31:10 -04:00
LogalDeveloper ee73e36f8e Expanded linting rules, added codespell and pip-audit, and fixed all violations.
CI / Formatting (push) Successful in 11s
CI / Linting (push) Successful in 11s
CI / Tests (push) Successful in 15s
CI / Type Checking (push) Successful in 21s
CI / Spelling (push) Successful in 12s
Dependency Audit / Dependency Audit (push) Successful in 7s
2026-02-20 10:30:36 -05:00
LogalDeveloper 0f2f43a55a Moved inline listener imports to module level in bot.py.
CI / Formatting (push) Successful in 11s
CI / Linting (push) Successful in 10s
CI / Tests (push) Successful in 14s
CI / Type Checking (push) Successful in 24s
2026-02-18 15:51:47 -05:00
LogalDeveloper fb92bc6766 Added user ID tracking to content tables and ingest-only CLI mode.
CI / Formatting (push) Successful in 10s
CI / Linting (push) Successful in 10s
CI / Tests (push) Successful in 15s
CI / Type Checking (push) Successful in 21s
2026-02-18 15:08:01 -05:00
57 changed files with 8816 additions and 1464 deletions
+35
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
+171
View File
@@ -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.
-128
View File
@@ -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()
-153
View File
@@ -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()
-326
View File
@@ -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.")
View File
-97
View File
@@ -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))
-79
View File
@@ -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))
-152
View File
@@ -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
-162
View File
@@ -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)
View File
+60 -26
View File
@@ -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
+20
View File
@@ -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())
+361
View File
@@ -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,
)
+122
View File
@@ -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()
+402
View File
@@ -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
+826
View File
@@ -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,),
)
+10 -15
View File
@@ -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__
+225
View File
@@ -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))
+111
View File
@@ -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.
""" """
+215
View File
@@ -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)
+233
View File
@@ -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()
+257
View File
@@ -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.")
+15
View File
@@ -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,
)
+15
View File
@@ -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."""
+34
View File
@@ -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)
+15
View File
@@ -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."""
+15
View File
@@ -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"
+15
View File
@@ -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."""
+35
View File
@@ -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
+15
View File
@@ -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."""
+133
View File
@@ -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 == []
+470
View File
@@ -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) == []
+15
View File
@@ -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."""
+196
View File
@@ -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()
+15
View File
@@ -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."""
+129
View File
@@ -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()
+629
View File
@@ -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
+164
View File
@@ -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
)
+470
View File
@@ -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.")]
+74
View File
@@ -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"
+26
View File
@@ -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__
Generated
+714 -231
View File
File diff suppressed because it is too large Load Diff