Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 11 additions & 10 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -42,16 +42,17 @@ To set up and deploy the Discord bot, follow these steps:
uv sync --dev
```

3. add the following environment variables.

| Variable | Description | Default |
|----------------|----------------------------|------------|
| BOT_NAME | The name of the bot | "Hackster" |
| BOT_TOKEN | The token of the bot | *Required |
| CHANNEL_DEVLOG | The devlog channel id | 0 |
| DEBUG | Toggles debug mode | False |
| DEV_GUILD_IDS | The dev servers of the bot | [] |
| GUILD_IDS | The servers of the bot | *Required |
3. add the following environment variables. Grouped settings use a double underscore between the group and
the field, for example `BOT__TOKEN`, `CHANNEL__DEVLOG` and `ROLE__VERIFIED`. See `.test.env` for the full list.

| Variable | Description | Default |
|-----------------|----------------------------|------------|
| BOT__NAME | The name of the bot | "Hackster" |
| BOT__TOKEN | The token of the bot | *Required |
| CHANNEL__DEVLOG | The devlog channel id | 0 |
| DEBUG | Toggles debug mode | False |
| DEV_GUILD_IDS | The dev servers of the bot | [] |
| GUILD_IDS | The servers of the bot | *Required |

4. Now you are done! You can run the project using

Expand Down
15 changes: 4 additions & 11 deletions alembic/env.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,8 @@
import os
from logging.config import fileConfig

import dotenv
from alembic import context
from sqlalchemy import create_engine

dotenv.load_dotenv()

# this is the Alembic Config object, which provides
# access to the values within the .ini file in use.
config = context.config
Expand All @@ -31,13 +27,10 @@


def get_url() -> str:
user = os.getenv("MYSQL_USER", "noahbot")
password = os.getenv("MYSQL_PASSWORD", None)
server = os.getenv("MYSQL_HOST", "localhost")
port = os.getenv("MYSQL_PORT", "3306")
db = os.getenv("MYSQL_DATABASE", "noahbot_dev")
url = f"mariadb+pymysql://{user}:{password}@{server}:{port}/{db}?charset=utf8mb4"
return url
"""Sync MariaDB URL from ``DatabaseSettings``, not the legacy ``MYSQL_*`` vars."""
from src.core.config import settings

return settings.database.assemble_db_connection(async_=False)


def run_migrations_offline() -> None:
Expand Down
13 changes: 10 additions & 3 deletions scripts/seed_dynamic_roles.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

from dotenv import dotenv_values
from sqlalchemy.dialects.mysql import insert
from sqlalchemy.ext.asyncio import AsyncSession

# Add project root to path
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
Expand Down Expand Up @@ -114,7 +115,12 @@
]


async def _upsert_role(session, values: dict) -> None:
def role_env_var(env_suffix: str) -> str:
"""Return the nested ``ROLE__`` env var for a seed suffix."""
return f"ROLE__{env_suffix}"


async def _upsert_role(session: AsyncSession, values: dict) -> None:
"""Insert or update a dynamic role using MariaDB upsert."""
stmt = insert(DynamicRole).values(values)
# On duplicate key, update all fields except the unique key (key, category)
Expand All @@ -129,6 +135,7 @@ async def _upsert_role(session, values: dict) -> None:


async def seed(env_file: str) -> None:
"""Upsert dynamic roles from ``ROLE__`` variables in ``env_file``."""
env_values = dotenv_values(env_file)
logger.info(f"Loaded env from {env_file} ({len(env_values)} values)")

Expand All @@ -138,7 +145,7 @@ async def seed(env_file: str) -> None:
async with AsyncSessionLocal() as session:
# Seed standard dynamic roles
for env_suffix, category, key, display_name, extra in SEED_DATA:
env_var = f"ROLE_{env_suffix}"
env_var = role_env_var(env_suffix)
role_id_str = env_values.get(env_var)
if not role_id_str:
logger.warning(f"Skipping {env_var}: not found in {env_file}")
Expand All @@ -161,7 +168,7 @@ async def seed(env_file: str) -> None:

# Seed joinable roles
for env_suffix, key, display_name, description in JOINABLE_SEED_DATA:
env_var = f"ROLE_{env_suffix}"
env_var = role_env_var(env_suffix)
role_id_str = env_values.get(env_var)
if not role_id_str:
logger.warning(f"Skipping joinable {env_var}/{key}: not found in {env_file}")
Expand Down
40 changes: 24 additions & 16 deletions src/core/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,9 +60,16 @@ class DatabaseSettings(BaseModel):
CHARSET: str = "utf8mb4"
ASYNC: bool | None = None

def assemble_db_connection(self) -> str:
def assemble_db_connection(self, async_: bool | None = None) -> str:
"""Build a SQLAlchemy MariaDB URL.

``async_`` overrides ``ASYNC``. Alembic passes ``False`` because it needs
a sync driver. When neither is set, the async driver is used.
"""
use_async = self.ASYNC if async_ is None else async_
driver = "pymysql" if use_async is False else "asyncmy"
return (
f"mariadb+asyncmy://{self.USER}:{self.PASSWORD}@{self.HOST}:{self.PORT}/"
f"mariadb+{driver}://{self.USER}:{self.PASSWORD}@{self.HOST}:{self.PORT}/"
f"{self.DATABASE}?charset={self.CHARSET}"
)

Expand Down Expand Up @@ -132,15 +139,8 @@ class RolesSettings(BaseModel):
CHALLENGE_CREATOR: int | None = None
BOX_CREATOR: int | None = None
SHERLOCK_CREATOR: int | None = None
APRIL_ROLE_1: int | None = None
APRIL_ROLE_2: int | None = None
RANK_FIVE: int | None = None
RANK_TWENTY_FIVE: int | None = None
RANK_FIFTY: int | None = None
RANK_HUNDRED: int | None = None
RANK_ONE: int | None = None
RANK_TEN: int | None = None
ACADEMY_CBBH: int | None = None
SEASON_HOLO: int | None = None
SEASON_PLATINUM: int | None = None
SEASON_RUBY: int | None = None
Expand Down Expand Up @@ -201,8 +201,6 @@ class Global(BaseSettings):
HTB_API_KEY: str
guild_ids: list[int]
dev_guild_ids: list[int] = Field(default_factory=list)
APRIL_FLAG_1: str = ""
APRIL_FLAG_2: str = ""

SENTRY_DSN: str | None = None
LOG_LEVEL: str | int = "INFO"
Expand All @@ -216,9 +214,7 @@ class Global(BaseSettings):
WEBHOOK_TOKEN: str = ""

SLACK_FEEDBACK_WEBHOOK: str = ""
SLACK_WEBHOOK: str = ""
JIRA_WEBHOOK: str = ""
JIRA_SPOILER_WEBHOOK: str = ""

# Feedback service ingest (POST /api/ingest/discord)
FEEDBACK_SERVICE_URL: str = ""
Expand Down Expand Up @@ -262,6 +258,18 @@ def role_groups(self) -> dict[str, list[int]]:
}


settings = Global(
_env_file=os.environ.get("ENV_PATH") if os.environ.get("BOT_ENVIRONMENT") else ".test.env"
)
def resolve_env_file() -> str | None:
"""Choose the env file used to build settings.

Deployed processes set ``APP_ENV_FILE``. ``BOT_ENVIRONMENT`` is still accepted
so existing Vault configs keep booting. Either flag loads ``ENV_PATH`` when
it is set, and otherwise uses the process environment. Local runs load
``.test.env``.
"""
deployed = os.environ.get("APP_ENV_FILE") or os.environ.get("BOT_ENVIRONMENT")
if deployed:
return os.environ.get("ENV_PATH")
return ".test.env"


settings = Global(_env_file=resolve_env_file())
58 changes: 56 additions & 2 deletions tests/src/core/test_config.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,11 @@
import os
import unittest
from unittest.mock import patch

from pydantic import ValidationError

from src.core.config import Global
from scripts.seed_dynamic_roles import role_env_var
from src.core.config import Global, resolve_env_file
from src.core import settings


Expand Down Expand Up @@ -45,7 +48,7 @@ def minimal_settings(**overrides):
"dev_guild_ids": [7764771731239076051],
}
payload.update(overrides)
return Global(**payload)
return Global(_env_file=None, **payload)

def test_guild_ids_accept_string_snowflakes(self):
"""Test that guild IDs can be provided as digit strings and are coerced to ints."""
Expand Down Expand Up @@ -100,3 +103,54 @@ def test_dynamic_roles_are_optional(self):
def test_season_id_loads_from_nested_env_config(self):
"""Test that SEASON_ID is loaded under the new nested env contract."""
self.assertEqual(settings.SEASON_ID, 1)

def test_assemble_db_connection_defaults_to_async_driver(self):
"""Test that an unset ASYNC flag uses the async MariaDB driver."""
config = self.minimal_settings()
url = config.database.assemble_db_connection()
self.assertTrue(url.startswith("mariadb+asyncmy://bot:secret@localhost:3306/bot"))

def test_assemble_db_connection_uses_sync_driver_when_disabled(self):
"""Test that ASYNC false, and an explicit override, select the sync driver."""
database = {
"HOST": "localhost",
"PORT": 3306,
"DATABASE": "bot",
"USER": "bot",
"PASSWORD": "secret",
"ASYNC": False,
}
config = self.minimal_settings(database=database)
self.assertTrue(config.database.assemble_db_connection().startswith("mariadb+pymysql://"))

async_config = self.minimal_settings()
sync_url = async_config.database.assemble_db_connection(async_=False)
self.assertTrue(sync_url.startswith("mariadb+pymysql://"))

def test_resolve_env_file_defaults_to_test_env(self):
"""Test that local runs load .test.env unless a deploy flag is set."""
with patch.dict(os.environ, {"APP_ENV_FILE": "", "BOT_ENVIRONMENT": "", "ENV_PATH": ".env"}):
self.assertEqual(resolve_env_file(), ".test.env")

def test_resolve_env_file_prefers_app_env_file(self):
"""Test that APP_ENV_FILE selects ENV_PATH and ignores the legacy flag's absence."""
with patch.dict(
os.environ,
{"APP_ENV_FILE": "production", "BOT_ENVIRONMENT": "", "ENV_PATH": "/vault/secrets/.env"},
):
self.assertEqual(resolve_env_file(), "/vault/secrets/.env")

def test_resolve_env_file_accepts_legacy_bot_environment(self):
"""Test that BOT_ENVIRONMENT still selects ENV_PATH until Vault is renamed."""
with patch.dict(os.environ, {"APP_ENV_FILE": "", "BOT_ENVIRONMENT": "production", "ENV_PATH": ".env"}):
self.assertEqual(resolve_env_file(), ".env")

def test_resolve_env_file_without_env_path_uses_process_environment(self):
"""Test that a deploy flag with no ENV_PATH does not force .test.env."""
with patch.dict(os.environ, {"APP_ENV_FILE": "production", "BOT_ENVIRONMENT": ""}):
os.environ.pop("ENV_PATH", None)
self.assertIsNone(resolve_env_file())

def test_seed_role_env_var_uses_nested_delimiter(self):
"""Test that dynamic-role seeding reads ROLE__ vars, not ROLE_."""
self.assertEqual(role_env_var("RANK_ONE"), "ROLE__RANK_ONE")
Loading