diff --git a/README.md b/README.md index 2c17b8b3..d23b222b 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/alembic/env.py b/alembic/env.py index daef157d..97f1c42f 100644 --- a/alembic/env.py +++ b/alembic/env.py @@ -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 @@ -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: diff --git a/scripts/seed_dynamic_roles.py b/scripts/seed_dynamic_roles.py index ecc24b30..ee9c5569 100644 --- a/scripts/seed_dynamic_roles.py +++ b/scripts/seed_dynamic_roles.py @@ -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__)))) @@ -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) @@ -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)") @@ -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}") @@ -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}") diff --git a/src/core/config.py b/src/core/config.py index b468ebcf..b8819041 100644 --- a/src/core/config.py +++ b/src/core/config.py @@ -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}" ) @@ -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 @@ -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" @@ -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 = "" @@ -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()) diff --git a/tests/src/core/test_config.py b/tests/src/core/test_config.py index 0260baca..8181c105 100644 --- a/tests/src/core/test_config.py +++ b/tests/src/core/test_config.py @@ -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 @@ -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.""" @@ -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")